use core::{
fmt,
mem::{align_of, size_of},
};
use crate::utils::{mem::transmute, Init};
#[derive(Copy, Clone)]
#[repr(transparent)]
pub struct ClosureEnv(Option<&'static ()>);
impl const Default for ClosureEnv {
#[inline]
fn default() -> Self {
Self::INIT
}
}
impl Init for ClosureEnv {
const INIT: Self = Self(None);
}
impl fmt::Debug for ClosureEnv {
#[inline]
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("ClosureEnv")
}
}
#[derive(Debug, Copy, Clone)]
pub struct Closure {
func: unsafe extern "C" fn(ClosureEnv),
env: ClosureEnv,
}
impl Init for Closure {
const INIT: Closure = (|| {}).into_closure_const();
}
impl Default for Closure {
#[inline]
fn default() -> Self {
Self::INIT
}
}
impl Closure {
#[inline]
pub const unsafe fn from_raw_parts(
func: unsafe extern "C" fn(ClosureEnv),
env: ClosureEnv,
) -> Self {
Self { func, env }
}
pub const fn from_fn_const<T: FnOnce() + Copy + Send + 'static>(func: T) -> Self {
let size = size_of::<T>();
let align = align_of::<T>();
unsafe {
if size == 0 {
Self::from_raw_parts(trampoline_zst::<T>, ClosureEnv(None))
} else {
let env = core::intrinsics::const_allocate(size, align);
assert!(
!env.guaranteed_eq(core::ptr::null_mut()).unwrap_or(false),
"heap allocation failed"
);
env.cast::<T>().write(func);
Self::from_raw_parts(trampoline_indirect::<T>, transmute(env))
}
}
}
#[inline]
pub fn call(self) {
unsafe { (self.func)(self.env) }
}
#[inline]
pub const fn func(self) -> unsafe extern "C" fn(ClosureEnv) {
self.func
}
#[inline]
pub const fn env(self) -> ClosureEnv {
self.env
}
#[inline]
pub const fn as_raw_parts(self) -> (unsafe extern "C" fn(ClosureEnv), ClosureEnv) {
(self.func, self.env)
}
}
#[inline]
unsafe extern "C" fn trampoline_zst<T: FnOnce()>(_: ClosureEnv) {
let func: T = unsafe { transmute(()) };
func()
}
#[inline]
unsafe extern "C" fn trampoline_indirect<T: FnOnce()>(env: ClosureEnv) {
let p_func: *const T = unsafe { transmute(env) };
let func: T = unsafe { p_func.read() };
func()
}
#[const_trait]
pub trait IntoClosureConst {
fn into_closure_const(self) -> Closure;
}
impl const IntoClosureConst for Closure {
fn into_closure_const(self) -> Closure {
self
}
}
impl<T: FnOnce() + Copy + Send + 'static> const IntoClosureConst for T {
fn into_closure_const(self) -> Closure {
Closure::from_fn_const(self)
}
}
impl<T: FnOnce(&'static P0) + Copy + Send + 'static, P0: Sync + 'static> const IntoClosureConst
for (&'static P0, T)
{
fn into_closure_const(self) -> Closure {
#[inline]
unsafe extern "C" fn trampoline_ptr_spec<T: FnOnce(&'static P0), P0: 'static>(
env: ClosureEnv,
) {
let p0: &'static P0 = unsafe { transmute(env) };
let func: T = unsafe { transmute(()) };
func(p0)
}
if size_of::<T>() == 0 {
unsafe { Closure::from_raw_parts(trampoline_ptr_spec::<T, P0>, transmute(self.0)) }
} else {
(move || (self.1)(self.0)).into_closure_const()
}
}
}
impl<T: FnOnce(usize) + Copy + Send + 'static> const IntoClosureConst for (usize, T) {
fn into_closure_const(self) -> Closure {
#[inline]
unsafe extern "C" fn trampoline_usize_spec<T: FnOnce(usize)>(env: ClosureEnv) {
let p0: usize = unsafe { transmute(env) };
let func: T = unsafe { transmute(()) };
func(p0)
}
if size_of::<T>() == 0 {
unsafe { Closure::from_raw_parts(trampoline_usize_spec::<T>, transmute(self.0)) }
} else {
(move || (self.1)(self.0)).into_closure_const()
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use core::sync::atomic::{AtomicUsize, Ordering};
#[test]
fn nested() {
static STATE: AtomicUsize = AtomicUsize::new(0);
const C1: Closure = {
let value = 0x1234;
(move || {
STATE.fetch_add(value, Ordering::Relaxed);
})
.into_closure_const()
};
const C2: Closure = {
let c = C1;
(move || {
c.call();
c.call();
})
.into_closure_const()
};
const C3: Closure = {
let c = C2;
(move || {
c.call();
c.call();
})
.into_closure_const()
};
const C4: Closure = {
let c = C3;
(move || {
c.call();
c.call();
})
.into_closure_const()
};
STATE.store(0, Ordering::Relaxed);
C4.call();
assert_eq!(STATE.load(Ordering::Relaxed), 0x1234 * 8);
}
#[test]
fn same_fn_different_env() {
static STATE: AtomicUsize = AtomicUsize::new(0);
const fn adder(x: usize) -> impl FnOnce() + Copy + Send {
move || {
STATE.fetch_add(x, Ordering::Relaxed);
}
}
const ADD1: Closure = adder(1).into_closure_const();
const ADD2: Closure = adder(2).into_closure_const();
const ADD4: Closure = adder(4).into_closure_const();
STATE.store(0, Ordering::Relaxed);
ADD1.call();
assert_eq!(STATE.load(Ordering::Relaxed), 1);
ADD4.call();
assert_eq!(STATE.load(Ordering::Relaxed), 1 + 4);
ADD2.call();
assert_eq!(STATE.load(Ordering::Relaxed), 1 + 4 + 2);
}
#[test]
fn ptr_env_spec() {
const C: Closure = (&42, |x: &i32| assert_eq!(*x, 42)).into_closure_const();
C.call();
}
#[test]
fn usize_env_spec() {
const C: Closure = (42usize, |x: usize| assert_eq!(x, 42)).into_closure_const();
C.call();
}
}