Skip to main content

generic_atomics/
atomic_float_wrapper.rs

1use atomic_traits::{
2    fetch::{And, Nand, Or, Update, Xor},
3    AsPtr, Atomic, Bitwise, FromPtr,
4};
5use core::{
6    cell::UnsafeCell,
7    ops::Not,
8    sync::atomic::{AtomicU32, AtomicU64, Ordering},
9};
10use num_primitive::PrimitiveFloat;
11use num_traits::FromPrimitive;
12
13pub trait BitsAssociatedAtomic {
14    type F: PrimitiveFloat;
15    type AU: Atomic<Type = <Self::F as PrimitiveFloat>::Bits>
16        + Bitwise
17        + Update<Type = <Self::F as PrimitiveFloat>::Bits>
18        + AsPtr
19        + FromPtr;
20}
21
22impl BitsAssociatedAtomic for u32 {
23    type AU = AtomicU32;
24    type F = f32;
25}
26
27impl BitsAssociatedAtomic for u64 {
28    type AU = AtomicU64;
29    type F = f64;
30}
31
32#[repr(transparent)]
33pub struct AtomicFloatWrapper<F: PrimitiveFloat>(UnsafeCell<F>);
34
35// SAFETY: We only ever access the underlying data by refcasting to AtomicU32,
36// which guarantees no data races.
37#[expect(unsafe_code)]
38unsafe impl<F: PrimitiveFloat> Send for AtomicFloatWrapper<F> {}
39#[expect(unsafe_code)]
40unsafe impl<F: PrimitiveFloat> Sync for AtomicFloatWrapper<F> {}
41
42// Static assertions that the layout is identical, we cite these in a safety
43// comment in `AtomicF32::atom()`. Note that the alignment check is stricter
44// than we need, as it would still be safe if `AtomicU32` is less strictly-
45// aligned than our `f32`. Unlike with `AtomicF64`, this is unlikely to occur.
46const _: [(); core::mem::size_of::<AtomicU32>()] = [(); core::mem::size_of::<UnsafeCell<f32>>()];
47const _: [(); core::mem::align_of::<AtomicU32>()] = [(); core::mem::align_of::<UnsafeCell<f32>>()];
48
49impl<F, U> AtomicFloatWrapper<F>
50where
51    F: PrimitiveFloat<Bits = U> + FromPrimitive,
52    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
53{
54    #[inline]
55    pub const fn new(float: F) -> Self {
56        Self(UnsafeCell::new(float))
57    }
58
59    #[inline]
60    fn as_atomic_bits(&self) -> &<F::Bits as BitsAssociatedAtomic>::AU {
61        // Safety: All potentially shared reads/writes go through this, and the
62        // static assertions above ensure that AtomicU32 and UnsafeCell<f32> are
63        // compatible as pointers.
64        let ptr_inner = &raw const self.0;
65        let cast_ptr = ptr_inner.cast::<<F::Bits as BitsAssociatedAtomic>::AU>();
66        #[expect(unsafe_code)]
67        unsafe {
68            &*cast_ptr
69        }
70    }
71}
72
73pub trait Abs {
74    type Type;
75    fn fetch_abs(&self, order: Ordering) -> Self::Type;
76}
77
78impl<F, U> Abs for AtomicFloatWrapper<F>
79where
80    F: PrimitiveFloat<Bits = U> + FromPrimitive,
81    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
82{
83    type Type = F;
84
85    #[inline]
86    fn fetch_abs(&self, order: Ordering) -> F {
87        // nu = 0x7fff_ffff
88        let mz: F = F::from_f32(-0f32).expect("infallible");
89        let nu = mz.to_bits().not();
90        let value = self.as_atomic_bits().fetch_and(nu, order);
91        F::from_bits(value)
92    }
93}
94
95pub trait Neg {
96    type Type;
97    fn fetch_neg(&self, order: Ordering) -> Self::Type;
98}
99
100impl<F, U> Neg for AtomicFloatWrapper<F>
101where
102    F: PrimitiveFloat<Bits = U> + FromPrimitive,
103    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
104{
105    type Type = F;
106
107    #[inline]
108    fn fetch_neg(&self, order: Ordering) -> F {
109        // u  = 0x80000000
110        let mz: F = F::from_f32(-0f32).expect("infallible");
111        let u: U = mz.to_bits();
112
113        F::from_bits(self.as_atomic_bits().fetch_xor(u, order))
114    }
115}
116
117trait UseAtomicUpdate: atomic_traits::fetch::Update {
118    fn update_with<F>(&self, order: Ordering, update: F) -> Self::Type
119    where
120        F: FnMut(Self::Type) -> Self::Type;
121}
122
123impl<F, U> UseAtomicUpdate for AtomicFloatWrapper<F>
124where
125    F: PrimitiveFloat<Bits = U> + FromPrimitive,
126    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
127{
128    #[inline]
129    fn update_with<Fun>(&self, order: Ordering, mut update: Fun) -> Self::Type
130    where
131        Fun: FnMut(Self::Type) -> Self::Type,
132    {
133        self.fetch_update(order, fail_order_for(order), |f| Some(update(f)))
134            .expect("infallible")
135    }
136}
137
138pub trait UseUpdateToAdd {
139    type Type;
140    fn fetch_add_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type;
141}
142
143impl<F, U> UseUpdateToAdd for AtomicFloatWrapper<F>
144where
145    F: PrimitiveFloat<Bits = U> + FromPrimitive,
146    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
147{
148    type Type = F;
149
150    #[inline(always)]
151    fn fetch_add_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type {
152        self.update_with(order, |f| f + val)
153    }
154}
155
156pub trait UseUpdateToSub {
157    type Type;
158    fn fetch_sub_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type;
159}
160
161impl<F, U> UseUpdateToSub for AtomicFloatWrapper<F>
162where
163    F: PrimitiveFloat<Bits = U> + FromPrimitive,
164    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
165{
166    type Type = F;
167
168    #[inline(always)]
169    fn fetch_sub_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type {
170        self.update_with(order, |f| f - val)
171    }
172}
173
174pub trait UseUpdateToMin {
175    type Type;
176    fn fetch_min_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type;
177}
178
179impl<F, U> UseUpdateToMin for AtomicFloatWrapper<F>
180where
181    F: PrimitiveFloat<Bits = U> + FromPrimitive,
182    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
183{
184    type Type = F;
185
186    #[inline(always)]
187    fn fetch_min_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type {
188        self.update_with(order, |f| f.min(val))
189    }
190}
191
192pub trait UseUpdateToMax {
193    type Type;
194    fn fetch_max_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type;
195}
196
197impl<F, U> UseUpdateToMax for AtomicFloatWrapper<F>
198where
199    F: PrimitiveFloat<Bits = U> + FromPrimitive,
200    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
201{
202    type Type = F;
203
204    #[inline(always)]
205    fn fetch_max_via_update(&self, val: Self::Type, order: Ordering) -> Self::Type {
206        self.update_with(order, |f| f.max(val))
207    }
208}
209
210impl<F, U> Atomic for AtomicFloatWrapper<F>
211where
212    F: PrimitiveFloat<Bits = U> + FromPrimitive,
213    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
214{
215    type Type = F;
216
217    #[inline]
218    fn new(float: F) -> Self {
219        Self(UnsafeCell::new(float))
220    }
221
222    #[inline]
223    fn get_mut(&mut self) -> &mut F {
224        // SAFETY: the mutable reference guarantees unique ownership.
225        let get = self.0.get();
226        #[expect(unsafe_code)]
227        unsafe {
228            &mut *get
229        }
230    }
231
232    #[inline]
233    fn into_inner(self) -> F {
234        self.0.into_inner()
235    }
236
237    #[inline]
238    fn load(&self, ordering: Ordering) -> F {
239        let value = self.as_atomic_bits().load(ordering);
240        F::from_bits(value)
241    }
242
243    #[inline]
244    fn store(&self, value: F, ordering: Ordering) {
245        let val = value.to_bits();
246        self.as_atomic_bits().store(val, ordering);
247    }
248
249    #[inline]
250    fn swap(&self, new_value: F, ordering: Ordering) -> F {
251        F::from_bits(self.as_atomic_bits().swap(new_value.to_bits(), ordering))
252    }
253
254    #[inline]
255    #[expect(deprecated)]
256    fn compare_and_swap(&self, current: F, new: F, order: Ordering) -> F {
257        F::from_bits(self.as_atomic_bits().compare_and_swap(
258            current.to_bits(),
259            new.to_bits(),
260            order,
261        ))
262    }
263
264    #[inline]
265    fn compare_exchange(
266        &self,
267        current: F,
268        new: F,
269        success: Ordering,
270        failure: Ordering,
271    ) -> Result<F, F> {
272        let current1 = current.to_bits();
273        let bits = new.to_bits();
274        convert_result(
275            self.as_atomic_bits()
276                .compare_exchange(current1, bits, success, failure),
277        )
278    }
279
280    #[inline]
281    fn compare_exchange_weak(
282        &self,
283        current: F,
284        new: F,
285        success: Ordering,
286        failure: Ordering,
287    ) -> Result<F, F> {
288        convert_result(self.as_atomic_bits().compare_exchange_weak(
289            current.to_bits(),
290            new.to_bits(),
291            success,
292            failure,
293        ))
294    }
295}
296
297impl<F, U> Update for AtomicFloatWrapper<F>
298where
299    F: PrimitiveFloat<Bits = U> + FromPrimitive,
300    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
301{
302    type Type = F;
303
304    #[inline]
305    fn fetch_update<Fun>(
306        &self,
307        set_order: Ordering,
308        fetch_order: Ordering,
309        mut update: Fun,
310    ) -> Result<F, F>
311    where
312        Fun: FnMut(F) -> Option<F>,
313    {
314        let atomic_bits = self.as_atomic_bits();
315        let res = U::AU::fetch_update(atomic_bits, set_order, fetch_order, |prev| {
316            update(F::from_bits(prev)).map(F::to_bits)
317        });
318        convert_result(res)
319    }
320}
321
322impl<F, U> Bitwise for AtomicFloatWrapper<F>
323where
324    F: PrimitiveFloat<Bits = U> + FromPrimitive,
325    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
326{
327}
328
329impl<F, U> And for AtomicFloatWrapper<F>
330where
331    F: PrimitiveFloat<Bits = U> + FromPrimitive,
332    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
333{
334    type Type = F;
335
336    #[inline]
337    fn fetch_and(&self, val: F, order: Ordering) -> F {
338        let val = F::to_bits(val);
339        let r = self.as_atomic_bits().fetch_and(val, order);
340        F::from_bits(r)
341    }
342}
343
344impl<F, U> Nand for AtomicFloatWrapper<F>
345where
346    F: PrimitiveFloat<Bits = U> + FromPrimitive,
347    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
348{
349    type Type = F;
350
351    #[inline]
352    fn fetch_nand(&self, val: F, order: Ordering) -> F {
353        let val = F::to_bits(val);
354        let r = self.as_atomic_bits().fetch_nand(val, order);
355        F::from_bits(r)
356    }
357}
358
359impl<F, U> Or for AtomicFloatWrapper<F>
360where
361    F: PrimitiveFloat<Bits = U> + FromPrimitive,
362    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
363{
364    type Type = F;
365
366    #[inline]
367    fn fetch_or(&self, val: F, order: Ordering) -> F {
368        let val = F::to_bits(val);
369        let r = self.as_atomic_bits().fetch_or(val, order);
370        F::from_bits(r)
371    }
372}
373
374impl<F, U> Xor for AtomicFloatWrapper<F>
375where
376    F: PrimitiveFloat<Bits = U> + FromPrimitive,
377    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
378{
379    type Type = F;
380
381    #[inline]
382    fn fetch_xor(&self, val: F, order: Ordering) -> F {
383        let val = F::to_bits(val);
384        let r = self.as_atomic_bits().fetch_xor(val, order);
385        F::from_bits(r)
386    }
387}
388
389impl<F, U> FromPtr for AtomicFloatWrapper<F>
390where
391    F: PrimitiveFloat<Bits = U> + FromPrimitive,
392    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
393{
394    #[inline(always)]
395    #[expect(unsafe_code)]
396    unsafe fn from_ptr<'a>(ptr: *mut F) -> &'a Self {
397        let ptr_u = ptr.cast::<U>();
398        let ref_atomic_u: &<U as BitsAssociatedAtomic>::AU = unsafe { U::AU::from_ptr(ptr_u) };
399        let ptr_atomic_u = core::ptr::from_ref::<<U as BitsAssociatedAtomic>::AU>(ref_atomic_u);
400        let ptr_t = ptr_atomic_u.cast::<Self>();
401        unsafe { &*ptr_t }
402    }
403}
404
405impl<F, U> AsPtr for AtomicFloatWrapper<F>
406where
407    F: PrimitiveFloat<Bits = U> + FromPrimitive,
408    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
409{
410    #[inline(always)]
411    fn as_ptr(&self) -> *mut F {
412        let src = self.as_atomic_bits().as_ptr();
413        src.cast::<F>()
414    }
415}
416
417impl<F, U> Default for AtomicFloatWrapper<F>
418where
419    F: PrimitiveFloat<Bits = U> + FromPrimitive,
420    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
421{
422    #[inline(always)]
423    fn default() -> Self {
424        Self(UnsafeCell::new(F::default()))
425    }
426}
427
428impl<F, U> core::fmt::Debug for AtomicFloatWrapper<F>
429where
430    F: PrimitiveFloat<Bits = U> + FromPrimitive,
431    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
432{
433    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
434        core::fmt::Debug::fmt(&self.load(Ordering::SeqCst), f)
435    }
436}
437
438impl<F, U> From<f32> for AtomicFloatWrapper<F>
439where
440    F: PrimitiveFloat<Bits = U> + FromPrimitive,
441    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
442{
443    #[inline]
444    fn from(f: f32) -> Self {
445        Self::new(F::from_f32(f).expect("infallible"))
446    }
447}
448
449impl<F, U> From<f64> for AtomicFloatWrapper<F>
450where
451    F: PrimitiveFloat<Bits = U> + FromPrimitive,
452    U: BitsAssociatedAtomic<F = F> + Not<Output = U>,
453{
454    #[inline(always)]
455    fn from(f: f64) -> Self {
456        Self::new(F::from_f64(f).expect("infallible"))
457    }
458}
459
460#[inline(always)]
461fn convert_result<F: PrimitiveFloat>(r: Result<F::Bits, F::Bits>) -> Result<F, F> {
462    r.map(F::from_bits).map_err(F::from_bits)
463}
464
465#[inline]
466fn fail_order_for(order: Ordering) -> Ordering {
467    match order {
468        Ordering::Release | Ordering::Relaxed => Ordering::Relaxed,
469        Ordering::Acquire | Ordering::AcqRel => Ordering::Acquire,
470        Ordering::SeqCst => Ordering::SeqCst,
471        o => unreachable!("Unknown ordering: {:?} (file a bug with atomic_float)", o),
472    }
473}