use super::placement;
use crate::topology::MAX_NUMA_NODE_IDS;
#[cfg(not(feature = "std"))]
extern crate alloc;
#[cfg(not(feature = "std"))]
use alloc::vec::Vec;
#[cfg(feature = "std")]
use std::vec::Vec;
use melinoe::sync::{sync_region_scope, SyncRegionToken};
pub struct SyncRegionPlacement<'brand> {
token: SyncRegionToken<'brand>,
}
fn assert_distinct_node_ids(nodes: &[crate::NumaNode]) {
let mut seen = [0u64; MAX_NUMA_NODE_IDS / u64::BITS as usize];
for node in nodes {
let index = node.id.index();
assert!(
index < MAX_NUMA_NODE_IDS,
"invariant: NUMA node id {index} exceeds the {MAX_NUMA_NODE_IDS}-node cap"
);
let (word, bit) = (index / u64::BITS as usize, index % u64::BITS as usize);
assert!(
seen[word] & (1 << bit) == 0,
"invariant: CpuTopology NUMA node ids must be pairwise distinct; \
duplicate id {index} would give two placement capabilities the same \
tag and let them alias one pinned cell"
);
seen[word] |= 1 << bit;
}
}
impl<'brand> SyncRegionPlacement<'brand> {
#[must_use]
#[inline]
pub const fn cell<T>(&self, value: T) -> melinoe::MelinoeCell<'brand, T> {
melinoe::MelinoeCell::new(value)
}
#[inline]
pub fn read<'a, T>(
&'a self,
cell: &'a melinoe::MelinoeCell<'brand, T>,
) -> melinoe::MelinoeRef<'a, 'brand, T> {
cell.borrow(&self.token)
}
#[inline]
pub fn write<'a, T>(
&'a mut self,
cell: &'a melinoe::MelinoeCell<'brand, T>,
) -> melinoe::MelinoeMut<'a, 'brand, T> {
cell.borrow_mut(&mut self.token)
}
#[must_use]
pub fn split(self, topology: &crate::CpuTopology) -> Vec<placement::NumaNodePlacement<'brand>> {
let nodes = topology.numa_nodes();
assert_distinct_node_ids(nodes);
let this = core::mem::ManuallyDrop::new(self);
let mut split = Vec::with_capacity(nodes.len());
for node in nodes {
unsafe {
split.push(placement::NumaNodePlacement {
node_id: node.id,
token: core::ptr::read(core::ptr::addr_of!(this.token)),
});
}
}
split
}
pub fn split_with<F, R>(self, topology: &crate::CpuTopology, f: F) -> R
where
F: FnOnce(&mut [placement::NumaNodePlacement<'brand>]) -> R,
{
const MAX_STACK_NODES: usize = 128;
let nodes = topology.numa_nodes();
assert_distinct_node_ids(nodes);
let this = core::mem::ManuallyDrop::new(self);
let num_nodes = nodes.len();
if num_nodes <= MAX_STACK_NODES {
struct DropGuard<'brand> {
ptr: *mut placement::NumaNodePlacement<'brand>,
initialized: usize,
}
impl Drop for DropGuard<'_> {
fn drop(&mut self) {
for i in 0..self.initialized {
unsafe {
core::ptr::drop_in_place(self.ptr.add(i));
}
}
}
}
let mut buf = core::mem::MaybeUninit::<
[placement::NumaNodePlacement<'brand>; MAX_STACK_NODES],
>::uninit();
let buf_ptr = buf
.as_mut_ptr()
.cast::<placement::NumaNodePlacement<'brand>>();
let mut guard = DropGuard {
ptr: buf_ptr,
initialized: 0,
};
for (i, node) in nodes.iter().enumerate() {
unsafe {
buf_ptr.add(i).write(placement::NumaNodePlacement {
node_id: node.id,
token: core::ptr::read(core::ptr::addr_of!(this.token)),
});
}
guard.initialized += 1;
}
let slice = unsafe { core::slice::from_raw_parts_mut(buf_ptr, num_nodes) };
f(slice)
} else {
let mut split = Vec::with_capacity(num_nodes);
for node in nodes {
unsafe {
split.push(placement::NumaNodePlacement {
node_id: node.id,
token: core::ptr::read(core::ptr::addr_of!(this.token)),
});
}
}
f(&mut split)
}
}
#[must_use]
#[inline]
pub fn split_static<const A: u32, const B: u32>(
self,
) -> (
placement::ConstNumaNodePlacement<'brand, A>,
placement::ConstNumaNodePlacement<'brand, B>,
) {
struct AssertDisjoint<const A: u32, const B: u32>;
impl<const A: u32, const B: u32> AssertDisjoint<A, B> {
const OK: () = {
assert!(A != B, "Static NUMA node split must be disjoint");
};
}
let () = AssertDisjoint::<A, B>::OK;
let this = core::mem::ManuallyDrop::new(self);
unsafe {
(
placement::ConstNumaNodePlacement {
token: core::ptr::read(core::ptr::addr_of!(this.token)),
},
placement::ConstNumaNodePlacement {
token: core::ptr::read(core::ptr::addr_of!(this.token)),
},
)
}
}
#[must_use]
#[inline]
pub fn split_static_3<const A: u32, const B: u32, const C: u32>(
self,
) -> (
placement::ConstNumaNodePlacement<'brand, A>,
placement::ConstNumaNodePlacement<'brand, B>,
placement::ConstNumaNodePlacement<'brand, C>,
) {
struct AssertDisjoint<const A: u32, const B: u32, const C: u32>;
impl<const A: u32, const B: u32, const C: u32> AssertDisjoint<A, B, C> {
const OK: () = {
assert!(
A != B && A != C && B != C,
"Static NUMA node split must be disjoint"
);
};
}
let () = AssertDisjoint::<A, B, C>::OK;
let this = core::mem::ManuallyDrop::new(self);
unsafe {
(
placement::ConstNumaNodePlacement {
token: core::ptr::read(core::ptr::addr_of!(this.token)),
},
placement::ConstNumaNodePlacement {
token: core::ptr::read(core::ptr::addr_of!(this.token)),
},
placement::ConstNumaNodePlacement {
token: core::ptr::read(core::ptr::addr_of!(this.token)),
},
)
}
}
#[must_use]
#[inline]
pub fn split_static_4<const A: u32, const B: u32, const C: u32, const D: u32>(
self,
) -> (
placement::ConstNumaNodePlacement<'brand, A>,
placement::ConstNumaNodePlacement<'brand, B>,
placement::ConstNumaNodePlacement<'brand, C>,
placement::ConstNumaNodePlacement<'brand, D>,
) {
struct AssertDisjoint<const A: u32, const B: u32, const C: u32, const D: u32>;
impl<const A: u32, const B: u32, const C: u32, const D: u32> AssertDisjoint<A, B, C, D> {
const OK: () = {
assert!(
A != B && A != C && A != D && B != C && B != D && C != D,
"Static NUMA node split must be disjoint"
);
};
}
let () = AssertDisjoint::<A, B, C, D>::OK;
let this = core::mem::ManuallyDrop::new(self);
unsafe {
(
placement::ConstNumaNodePlacement {
token: core::ptr::read(core::ptr::addr_of!(this.token)),
},
placement::ConstNumaNodePlacement {
token: core::ptr::read(core::ptr::addr_of!(this.token)),
},
placement::ConstNumaNodePlacement {
token: core::ptr::read(core::ptr::addr_of!(this.token)),
},
placement::ConstNumaNodePlacement {
token: core::ptr::read(core::ptr::addr_of!(this.token)),
},
)
}
}
#[must_use]
#[inline]
pub fn project_static<const NODE_ID: u32>(
self,
) -> placement::ConstNumaNodePlacement<'brand, NODE_ID> {
let this = core::mem::ManuallyDrop::new(self);
placement::ConstNumaNodePlacement {
token: unsafe { core::ptr::read(core::ptr::addr_of!(this.token)) },
}
}
}
#[inline]
pub fn sync_region_placement_scope<R>(
f: impl for<'brand> FnOnce(SyncRegionPlacement<'brand>) -> R,
) -> R {
sync_region_scope(|token| f(SyncRegionPlacement { token }))
}
#[cfg(test)]
mod tests {
use super::assert_distinct_node_ids;
use crate::{MemoryTier, NumaNode, NumaNodeId};
fn node(id: u32) -> NumaNode {
NumaNode {
id: NumaNodeId::new(id),
processors: Box::new([id]),
distances: Box::new([10]),
memory_tier: MemoryTier::Dram,
}
}
#[test]
fn distinct_node_ids_are_accepted() {
assert_distinct_node_ids(&[node(0), node(1), node(7), node(1023)]);
}
#[test]
#[should_panic(expected = "pairwise distinct")]
fn repeated_node_id_is_rejected() {
assert_distinct_node_ids(&[node(0), node(1), node(0)]);
}
#[test]
#[should_panic(expected = "exceeds the 1024-node cap")]
fn out_of_range_node_id_is_rejected() {
assert_distinct_node_ids(&[node(1024)]);
}
}