#![allow(clippy::missing_errors_doc)]
use core::fmt;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum Axis {
Inner,
Outer,
Consumer,
CounterShare,
Radius,
Gating,
Stride,
Segmented,
ContentPrefix,
TypeTag,
Version,
Async,
Bounds,
}
impl Axis {
pub const ALL: [Axis; 13] = [
Axis::Inner,
Axis::Outer,
Axis::Consumer,
Axis::CounterShare,
Axis::Radius,
Axis::Gating,
Axis::Stride,
Axis::Segmented,
Axis::ContentPrefix,
Axis::TypeTag,
Axis::Version,
Axis::Async,
Axis::Bounds,
];
#[inline(always)]
pub const fn bit(self) -> u16 {
match self {
Axis::Inner => 0,
Axis::Outer => 1,
Axis::Consumer => 2,
Axis::CounterShare => 3,
Axis::Radius => 4,
Axis::Gating => 5,
Axis::Stride => 6,
Axis::Segmented => 7,
Axis::ContentPrefix => 8,
Axis::TypeTag => 9,
Axis::Version => 10,
Axis::Async => 11,
Axis::Bounds => 12,
}
}
}
impl fmt::Display for Axis {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let name = match self {
Axis::Inner => "K_inner",
Axis::Outer => "K_outer",
Axis::Consumer => "K_consumer",
Axis::CounterShare => "K_counter_share",
Axis::Radius => "K_radius",
Axis::Gating => "K_gating",
Axis::Stride => "K_stride",
Axis::Segmented => "K_segmented",
Axis::ContentPrefix => "K_content_prefix",
Axis::TypeTag => "K_type_tag",
Axis::Version => "K_version",
Axis::Async => "K_async",
Axis::Bounds => "K_bounds",
};
f.write_str(name)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct AxisMask(u16);
impl AxisMask {
const VALID_BITS: u16 = (1u16 << 13) - 1;
pub const EMPTY: AxisMask = AxisMask(0);
pub const ALL: AxisMask = AxisMask(Self::VALID_BITS);
pub const fn from_axes(axes: &[Axis]) -> Self {
let mut bits = 0u16;
let mut i = 0;
while i < axes.len() {
bits |= 1u16 << axes[i].bit();
i += 1;
}
AxisMask(bits)
}
pub const fn from_bits(bits: u16) -> Self {
AxisMask(bits & Self::VALID_BITS)
}
#[inline(always)]
pub const fn bits(self) -> u16 {
self.0
}
#[inline(always)]
pub const fn contains(self, a: Axis) -> bool {
(self.0 >> a.bit()) & 1 == 1
}
#[inline(always)]
pub const fn count(self) -> u32 {
self.0.count_ones()
}
#[inline(always)]
pub const fn union(self, other: AxisMask) -> AxisMask {
AxisMask(self.0 | other.0)
}
#[inline(always)]
pub const fn intersection(self, other: AxisMask) -> AxisMask {
AxisMask(self.0 & other.0)
}
#[inline(always)]
pub const fn satisfies(self, other: AxisMask) -> bool {
(self.0 & other.0) == other.0
}
#[inline(always)]
pub const fn distance(self, other: AxisMask) -> u32 {
(self.0 ^ other.0).count_ones()
}
}
impl fmt::Display for AxisMask {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let mut first = true;
f.write_str("{")?;
for a in Axis::ALL {
if self.contains(a) {
if !first {
f.write_str(", ")?;
}
fmt::Display::fmt(&a, f)?;
first = false;
}
}
f.write_str("}")
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Fusion {
pub axes: AxisMask,
}
impl Fusion {
pub const fn pair(a: Axis, b: Axis) -> Self {
Self {
axes: AxisMask::from_axes(&[a, b]),
}
}
pub const fn triple(a: Axis, b: Axis, c: Axis) -> Self {
Self {
axes: AxisMask::from_axes(&[a, b, c]),
}
}
#[inline(always)]
pub const fn axis_count(self) -> u32 {
self.axes.count()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_signature_contains_nothing() {
assert!(!AxisMask::EMPTY.contains(Axis::Inner));
assert_eq!(AxisMask::EMPTY.count(), 0);
}
#[test]
fn all_signature_contains_everything() {
for a in Axis::ALL {
assert!(AxisMask::ALL.contains(a));
}
assert_eq!(AxisMask::ALL.count(), Axis::ALL.len() as u32);
}
#[test]
fn from_axes_packs_bits_correctly() {
let s = AxisMask::from_axes(&[Axis::Inner, Axis::Gating]);
assert!(s.contains(Axis::Inner));
assert!(s.contains(Axis::Gating));
assert!(!s.contains(Axis::Outer));
assert_eq!(s.count(), 2);
}
#[test]
fn satisfies_is_superset_check() {
let chase_lev = AxisMask::EMPTY;
let khl = AxisMask::from_axes(&[
Axis::Inner,
Axis::Outer,
Axis::CounterShare,
Axis::Radius,
Axis::Gating,
]);
let request_reply = AxisMask::EMPTY;
let producer_fast = AxisMask::from_axes(&[Axis::Inner, Axis::Outer]);
assert!(chase_lev.satisfies(request_reply));
assert!(khl.satisfies(request_reply));
assert!(khl.satisfies(producer_fast));
assert!(!chase_lev.satisfies(producer_fast));
}
#[test]
fn distance_counts_differing_axes() {
let a = AxisMask::from_axes(&[Axis::Inner, Axis::Outer]);
let b = AxisMask::from_axes(&[Axis::Inner, Axis::Gating]);
assert_eq!(a.distance(b), 2);
assert_eq!(a.distance(a), 0);
assert_eq!(
AxisMask::EMPTY.distance(AxisMask::ALL),
Axis::ALL.len() as u32
);
}
#[test]
fn pair_fusion_has_axis_count_two() {
let f = Fusion::pair(Axis::Inner, Axis::Gating);
assert_eq!(f.axis_count(), 2);
}
#[test]
fn triple_fusion_has_axis_count_three() {
let f = Fusion::triple(Axis::Radius, Axis::Inner, Axis::Gating);
assert_eq!(f.axis_count(), 3);
}
#[test]
fn display_renders_axis_names() {
let s = AxisMask::from_axes(&[Axis::Inner, Axis::Radius]);
let rendered = format!("{s}");
assert!(rendered.contains("K_inner"));
assert!(rendered.contains("K_radius"));
}
}