use atomic_traits::{
fetch::{And, Nand, Or, Update, Xor},
AsPtr, Atomic, Bitwise, FromPtr,
};
use core::{
cell::UnsafeCell,
ops::Not,
sync::atomic::{AtomicU32, AtomicU64, Ordering},
};
use num_primitive::PrimitiveFloat;
use num_traits::FromPrimitive;
pub trait BitsAssociatedAtomic {
type F: PrimitiveFloat;
type AU: Atomic<Type = <Self::F as PrimitiveFloat>::Bits>
+ Bitwise
+ Update<Type = <Self::F as PrimitiveFloat>::Bits>
+ AsPtr
+ FromPtr;
}
impl BitsAssociatedAtomic for u32 {
type AU = AtomicU32;
type F = f32;
}
impl BitsAssociatedAtomic for u64 {
type AU = AtomicU64;
type F = f64;
}
#[repr(transparent)]
pub struct AtomicFloatWrapper<F: PrimitiveFloat>(UnsafeCell<F>);
#[expect(unsafe_code)]
unsafe impl<F: PrimitiveFloat> Send for AtomicFloatWrapper<F> {}
#[expect(unsafe_code)]
unsafe impl<F: PrimitiveFloat> Sync for AtomicFloatWrapper<F> {}
const _: [(); core::mem::size_of::<AtomicU32>()] = [(); core::mem::size_of::<UnsafeCell<f32>>()];
const _: [(); core::mem::align_of::<AtomicU32>()] = [(); core::mem::align_of::<UnsafeCell<f32>>()];
impl<F, U> AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
#[inline]
pub const fn new(float: F) -> Self {
Self(UnsafeCell::new(float))
}
#[inline]
fn as_atomic_bits(&self) -> &<F::Bits as BitsAssociatedAtomic>::AU {
let ptr_inner = &raw const self.0;
let cast_ptr = ptr_inner.cast::<<F::Bits as BitsAssociatedAtomic>::AU>();
#[expect(unsafe_code)]
unsafe {
&*cast_ptr
}
}
}
pub trait Abs {
type Type;
fn fetch_abs(&self, order: Ordering) -> Self::Type;
}
impl<F, U> Abs for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
type Type = F;
#[inline]
fn fetch_abs(&self, order: Ordering) -> F {
let mz: F = F::from_f32(-0f32).expect("infallible");
let nu = mz.to_bits().not();
let value = self.as_atomic_bits().fetch_and(nu, order);
F::from_bits(value)
}
}
pub trait Neg {
type Type;
fn fetch_neg(&self, order: Ordering) -> Self::Type;
}
impl<F, U> Neg for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
type Type = F;
#[inline]
fn fetch_neg(&self, order: Ordering) -> F {
let mz: F = F::from_f32(-0f32).expect("infallible");
let u: U = mz.to_bits();
F::from_bits(self.as_atomic_bits().fetch_xor(u, order))
}
}
trait UseAtomicUpdate: atomic_traits::fetch::Update {
fn update_with<F>(&self, order: Ordering, update: F) -> Self::Type
where
F: FnMut(Self::Type) -> Self::Type;
}
impl<F, U> UseAtomicUpdate for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
#[inline]
fn update_with<Fun>(&self, order: Ordering, mut update: Fun) -> Self::Type
where
Fun: FnMut(Self::Type) -> Self::Type,
{
self.fetch_update(order, fail_order_for(order), |f| Some(update(f)))
.expect("infallible")
}
}
pub trait UseUpdateToAdd {
type Type;
fn fetch_add_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type;
}
impl<F, U> UseUpdateToAdd for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
type Type = F;
#[inline(always)]
fn fetch_add_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type {
self.update_with(order, |f| f + val)
}
}
pub trait UseUpdateToSub {
type Type;
fn fetch_sub_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type;
}
impl<F, U> UseUpdateToSub for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
type Type = F;
#[inline(always)]
fn fetch_sub_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type {
self.update_with(order, |f| f - val)
}
}
pub trait UseUpdateToMin {
type Type;
fn fetch_min_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type;
}
impl<F, U> UseUpdateToMin for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
type Type = F;
#[inline(always)]
fn fetch_min_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type {
self.update_with(order, |f| f.min(val))
}
}
pub trait UseUpdateToMax {
type Type;
fn fetch_max_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type;
}
impl<F, U> UseUpdateToMax for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
type Type = F;
#[inline(always)]
fn fetch_max_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type {
self.update_with(order, |f| f.max(val))
}
}
impl<F, U> Atomic for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
type Type = F;
#[inline]
fn new(float: F) -> Self {
Self(UnsafeCell::new(float))
}
#[inline]
fn get_mut(&mut self) -> &mut F {
let get = self.0.get();
#[expect(unsafe_code)]
unsafe {
&mut *get
}
}
#[inline]
fn into_inner(self) -> F {
self.0.into_inner()
}
#[inline]
fn load(&self, ordering: Ordering) -> F {
let value = self.as_atomic_bits().load(ordering);
F::from_bits(value)
}
#[inline]
fn store(&self, value: F, ordering: Ordering) {
let val = value.to_bits();
self.as_atomic_bits().store(val, ordering);
}
#[inline]
fn swap(&self, new_value: F, ordering: Ordering) -> F {
F::from_bits(self.as_atomic_bits().swap(new_value.to_bits(), ordering))
}
#[inline]
#[expect(deprecated)]
fn compare_and_swap(&self, current: F, new: F, order: Ordering) -> F {
F::from_bits(self.as_atomic_bits().compare_and_swap(
current.to_bits(),
new.to_bits(),
order,
))
}
#[inline]
fn compare_exchange(
&self,
current: F,
new: F,
success: Ordering,
failure: Ordering,
) -> Result<F, F> {
let current1 = current.to_bits();
let bits = new.to_bits();
convert_result(
self.as_atomic_bits()
.compare_exchange(current1, bits, success, failure),
)
}
#[inline]
fn compare_exchange_weak(
&self,
current: F,
new: F,
success: Ordering,
failure: Ordering,
) -> Result<F, F> {
convert_result(self.as_atomic_bits().compare_exchange_weak(
current.to_bits(),
new.to_bits(),
success,
failure,
))
}
}
impl<F, U> Update for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
type Type = F;
#[inline]
fn fetch_update<Fun>(
&self,
set_order: Ordering,
fetch_order: Ordering,
mut update: Fun,
) -> Result<F, F>
where
Fun: FnMut(F) -> Option<F>,
{
let atomic_bits = self.as_atomic_bits();
let res = U::AU::fetch_update(atomic_bits, set_order, fetch_order, |prev| {
update(F::from_bits(prev)).map(F::to_bits)
});
convert_result(res)
}
}
impl<F, U> Bitwise for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
}
impl<F, U> And for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
type Type = F;
#[inline]
fn fetch_and(&self, val: F, order: Ordering) -> F {
let val = F::to_bits(val);
let r = self.as_atomic_bits().fetch_and(val, order);
F::from_bits(r)
}
}
impl<F, U> Nand for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
type Type = F;
#[inline]
fn fetch_nand(&self, val: F, order: Ordering) -> F {
let val = F::to_bits(val);
let r = self.as_atomic_bits().fetch_nand(val, order);
F::from_bits(r)
}
}
impl<F, U> Or for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
type Type = F;
#[inline]
fn fetch_or(&self, val: F, order: Ordering) -> F {
let val = F::to_bits(val);
let r = self.as_atomic_bits().fetch_or(val, order);
F::from_bits(r)
}
}
impl<F, U> Xor for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
type Type = F;
#[inline]
fn fetch_xor(&self, val: F, order: Ordering) -> F {
let val = F::to_bits(val);
let r = self.as_atomic_bits().fetch_xor(val, order);
F::from_bits(r)
}
}
impl<F, U> FromPtr for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
#[inline(always)]
#[expect(unsafe_code)]
unsafe fn from_ptr<'a>(ptr: *mut F) -> &'a Self {
let ptr_u = ptr.cast::<U>();
let ref_atomic_u: &<U as BitsAssociatedAtomic>::AU = unsafe { U::AU::from_ptr(ptr_u) };
let ptr_atomic_u = core::ptr::from_ref::<<U as BitsAssociatedAtomic>::AU>(ref_atomic_u);
let ptr_t = ptr_atomic_u.cast::<Self>();
unsafe { &*ptr_t }
}
}
impl<F, U> AsPtr for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
#[inline(always)]
fn as_ptr(&self) -> *mut F {
let src = self.as_atomic_bits().as_ptr();
src.cast::<F>()
}
}
impl<F, U> Default for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
#[inline(always)]
fn default() -> Self {
Self(UnsafeCell::new(F::default()))
}
}
impl<F, U> core::fmt::Debug for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
core::fmt::Debug::fmt(&self.load(Ordering::SeqCst), f)
}
}
impl<F, U> From<f32> for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
#[inline]
fn from(f: f32) -> Self {
Self::new(F::from_f32(f).expect("infallible"))
}
}
impl<F, U> From<f64> for AtomicFloatWrapper<F>
where
F: PrimitiveFloat<Bits = U> + FromPrimitive,
U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
{
#[inline(always)]
fn from(f: f64) -> Self {
Self::new(F::from_f64(f).expect("infallible"))
}
}
#[inline(always)]
fn convert_result<F: PrimitiveFloat>(r: Result<F::Bits, F::Bits>) -> Result<F, F> {
r.map(F::from_bits).map_err(F::from_bits)
}
#[inline]
fn fail_order_for(order: Ordering) -> Ordering {
match order {
Ordering::Release | Ordering::Relaxed => Ordering::Relaxed,
Ordering::Acquire | Ordering::AcqRel => Ordering::Acquire,
Ordering::SeqCst => Ordering::SeqCst,
o => unreachable!("Unknown ordering: {:?} (file a bug with atomic_float)", o),
}
}