Skip to main content

cold_string/
arc.rs

1use alloc::{
2    borrow::{Cow, ToOwned},
3    boxed::Box,
4    str::Utf8Error,
5    string::String,
6};
7use core::{
8    cmp::Ordering,
9    fmt,
10    hash::{Hash, Hasher},
11    iter::FromIterator,
12    ops::Deref,
13    str,
14};
15
16#[cfg(all(loom, test, target_arch = "x86_64"))]
17use loom::sync::atomic::{fence, AtomicU16, AtomicU32, AtomicU8, AtomicUsize, Ordering::*};
18#[cfg(not(all(loom, test, target_arch = "x86_64")))]
19use portable_atomic::{fence, AtomicU16, AtomicU32, AtomicU8, AtomicUsize, Ordering::*};
20
21use crate::encoded::Encoded;
22
23#[doc(hidden)]
24pub trait RefCount: Send + Sync + 'static {
25    fn new() -> Self;
26    fn increment(&self);
27    fn decrement(&self) -> bool;
28
29    #[cfg(test)]
30    fn refs(&self) -> usize;
31
32    #[cfg(test)]
33    fn near_overflow() -> Self;
34
35    #[cfg(test)]
36    fn is_immortal(&self) -> bool;
37
38    #[cfg(all(test, not(all(loom, target_arch = "x86_64"))))]
39    fn min_immortal() -> Self;
40}
41
42macro_rules! impl_ref_count {
43    ($atomic:ty, $int:ty) => {
44        impl RefCount for $atomic {
45            #[inline]
46            fn new() -> Self {
47                <$atomic>::new(2)
48            }
49
50            #[inline]
51            fn increment(&self) {
52                let mut old = self.load(Relaxed);
53                loop {
54                    if old & 1 != 0 {
55                        return;
56                    }
57
58                    let new = if old == <$int>::MAX - 1 {
59                        <$int>::MAX
60                    } else {
61                        old + 2
62                    };
63                    match self.compare_exchange_weak(old, new, Relaxed, Relaxed) {
64                        Ok(_) => return,
65                        Err(actual) => old = actual,
66                    }
67                }
68            }
69
70            #[inline]
71            fn decrement(&self) -> bool {
72                self.fetch_sub(2, Release) == 2
73            }
74
75            #[cfg(test)]
76            fn refs(&self) -> usize {
77                (self.load(Relaxed) >> 1) as usize
78            }
79
80            #[cfg(test)]
81            fn near_overflow() -> Self {
82                Self::new(<$int>::MAX - 1)
83            }
84
85            #[cfg(test)]
86            fn is_immortal(&self) -> bool {
87                self.load(Relaxed) & 1 != 0
88            }
89
90            #[cfg(all(test, not(all(loom, target_arch = "x86_64"))))]
91            fn min_immortal() -> Self {
92                Self::new(1)
93            }
94        }
95    };
96}
97
98impl_ref_count!(AtomicU8, u8);
99impl_ref_count!(AtomicU16, u16);
100impl_ref_count!(AtomicU32, u32);
101impl_ref_count!(AtomicUsize, usize);
102
103#[doc(hidden)]
104#[repr(transparent)]
105pub struct ArcColdStringInner<A: RefCount> {
106    encoded: Encoded<A>,
107}
108
109/// A one-word, atomically reference-counted immutable UTF-8 string.
110///
111/// Strings up to one machine word are stored inline. Longer strings use one
112/// allocation containing the reference count, variable-length length, and bytes.
113/// If the reference count saturates, the allocation becomes immortal (and is
114/// never deallocated) rather than allowing the count to wrap.
115///
116/// ```
117/// use cold_string::ArcColdString;
118///
119/// let first = ArcColdString::new("a string longer than one machine word");
120/// let second = first.clone();
121/// assert_eq!(first, second);
122/// ```
123pub type ArcColdString = ArcColdStringInner<AtomicUsize>;
124
125/// An [`ArcColdString`] with an 8-bit reference count.
126///
127/// It supports up to 127 live references. Cloning beyond that makes the
128/// allocation immortal.
129pub type ArcColdString8 = ArcColdStringInner<AtomicU8>;
130
131/// An [`ArcColdString`] with a 16-bit reference count.
132///
133/// It supports up to 32,767 live references. Cloning beyond that makes the
134/// allocation immortal.
135pub type ArcColdString16 = ArcColdStringInner<AtomicU16>;
136
137/// An [`ArcColdString`] with a 32-bit reference count.
138///
139/// It supports up to 2,147,483,647 live references. Cloning beyond that makes
140/// the allocation immortal.
141pub type ArcColdString32 = ArcColdStringInner<AtomicU32>;
142
143impl<A: RefCount> ArcColdStringInner<A> {
144    pub fn from_utf8<B: AsRef<[u8]>>(bytes: B) -> Result<Self, Utf8Error> {
145        Ok(Self::new(str::from_utf8(bytes.as_ref())?))
146    }
147
148    /// # Safety
149    ///
150    /// `bytes` must contain valid UTF-8.
151    pub unsafe fn from_utf8_unchecked<B: AsRef<[u8]>>(bytes: B) -> Self {
152        Self::new(str::from_utf8_unchecked(bytes.as_ref()))
153    }
154
155    pub fn new<T: AsRef<str>>(value: T) -> Self {
156        let s = value.as_ref();
157        Self {
158            encoded: Encoded::new(s, A::new()),
159        }
160    }
161
162    #[rustversion::since(1.61)]
163    #[inline]
164    pub const fn new_inline_const(s: &str) -> Self {
165        Self {
166            encoded: Encoded::new_inline_const(s),
167        }
168    }
169
170    #[inline]
171    fn count(&self) -> &A {
172        debug_assert!(!self.is_inline());
173        unsafe { &(*self.encoded.heap_ptr().as_ptr()).header }
174    }
175
176    #[cfg(test)]
177    pub(crate) fn encoded_addr(&self) -> usize {
178        self.encoded.addr()
179    }
180
181    #[inline]
182    pub fn is_inline(&self) -> bool {
183        self.encoded.is_inline()
184    }
185
186    #[inline]
187    pub fn len(&self) -> usize {
188        self.encoded.len()
189    }
190
191    #[inline]
192    pub fn as_bytes(&self) -> &[u8] {
193        self.encoded.as_bytes()
194    }
195
196    #[inline]
197    pub fn as_str(&self) -> &str {
198        // SAFETY: constructors accept only valid UTF-8.
199        unsafe { str::from_utf8_unchecked(self.as_bytes()) }
200    }
201
202    #[inline]
203    pub fn is_empty(&self) -> bool {
204        self.len() == 0
205    }
206}
207
208impl<A: RefCount> Clone for ArcColdStringInner<A> {
209    #[inline]
210    fn clone(&self) -> Self {
211        if !self.is_inline() {
212            self.count().increment();
213        }
214
215        Self {
216            encoded: self.encoded,
217        }
218    }
219}
220
221impl<A: RefCount> Drop for ArcColdStringInner<A> {
222    #[inline]
223    fn drop(&mut self) {
224        if self.is_inline() {
225            return;
226        }
227
228        if !self.count().decrement() {
229            return;
230        }
231
232        fence(Acquire);
233        // SAFETY: this was the last reference and the count is synchronized.
234        unsafe { self.encoded.deallocate() }
235    }
236}
237
238impl<A: RefCount> Default for ArcColdStringInner<A> {
239    fn default() -> Self {
240        Self::new("")
241    }
242}
243
244impl<A: RefCount> Deref for ArcColdStringInner<A> {
245    type Target = str;
246
247    fn deref(&self) -> &str {
248        self.as_str()
249    }
250}
251
252impl<A: RefCount> PartialEq for ArcColdStringInner<A> {
253    fn eq(&self, other: &Self) -> bool {
254        self.encoded.addr() == other.encoded.addr() || self.as_bytes() == other.as_bytes()
255    }
256}
257
258impl<A: RefCount> Eq for ArcColdStringInner<A> {}
259
260impl<A: RefCount> Hash for ArcColdStringInner<A> {
261    fn hash<H: Hasher>(&self, state: &mut H) {
262        self.as_str().hash(state)
263    }
264}
265
266impl<A: RefCount> fmt::Debug for ArcColdStringInner<A> {
267    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
268        fmt::Debug::fmt(self.as_str(), f)
269    }
270}
271
272impl<A: RefCount> fmt::Display for ArcColdStringInner<A> {
273    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
274        fmt::Display::fmt(self.as_str(), f)
275    }
276}
277
278impl<A: RefCount> From<&str> for ArcColdStringInner<A> {
279    fn from(s: &str) -> Self {
280        Self::new(s)
281    }
282}
283
284impl<A: RefCount> From<String> for ArcColdStringInner<A> {
285    fn from(s: String) -> Self {
286        Self::new(&s)
287    }
288}
289
290impl<A: RefCount> From<Box<str>> for ArcColdStringInner<A> {
291    fn from(s: Box<str>) -> Self {
292        Self::new(&s)
293    }
294}
295
296impl<A: RefCount> From<ArcColdStringInner<A>> for String {
297    fn from(s: ArcColdStringInner<A>) -> Self {
298        s.as_str().to_owned()
299    }
300}
301
302impl<A: RefCount> From<ArcColdStringInner<A>> for Cow<'_, str> {
303    fn from(s: ArcColdStringInner<A>) -> Self {
304        Self::Owned(s.into())
305    }
306}
307
308impl<'a, A: RefCount> From<&'a ArcColdStringInner<A>> for Cow<'a, str> {
309    fn from(s: &'a ArcColdStringInner<A>) -> Self {
310        Self::Borrowed(s)
311    }
312}
313
314impl<'a, A: RefCount> From<Cow<'a, str>> for ArcColdStringInner<A> {
315    fn from(s: Cow<'a, str>) -> Self {
316        Self::new(s)
317    }
318}
319
320impl<A: RefCount> FromIterator<char> for ArcColdStringInner<A> {
321    fn from_iter<I: IntoIterator<Item = char>>(iter: I) -> Self {
322        Self::new(iter.into_iter().collect::<String>())
323    }
324}
325
326impl<A: RefCount> core::borrow::Borrow<str> for ArcColdStringInner<A> {
327    fn borrow(&self) -> &str {
328        self.as_str()
329    }
330}
331
332impl<A: RefCount> PartialEq<str> for ArcColdStringInner<A> {
333    fn eq(&self, other: &str) -> bool {
334        self.as_str() == other
335    }
336}
337
338impl<A: RefCount> PartialEq<ArcColdStringInner<A>> for str {
339    fn eq(&self, other: &ArcColdStringInner<A>) -> bool {
340        other == self
341    }
342}
343
344impl<A: RefCount> PartialEq<&str> for ArcColdStringInner<A> {
345    fn eq(&self, other: &&str) -> bool {
346        self == *other
347    }
348}
349
350impl<A: RefCount> PartialEq<ArcColdStringInner<A>> for &str {
351    fn eq(&self, other: &ArcColdStringInner<A>) -> bool {
352        other == *self
353    }
354}
355
356impl<A: RefCount> AsRef<str> for ArcColdStringInner<A> {
357    fn as_ref(&self) -> &str {
358        self.as_str()
359    }
360}
361
362impl<A: RefCount> AsRef<[u8]> for ArcColdStringInner<A> {
363    fn as_ref(&self) -> &[u8] {
364        self.as_bytes()
365    }
366}
367
368impl<A: RefCount> Ord for ArcColdStringInner<A> {
369    fn cmp(&self, other: &Self) -> Ordering {
370        self.as_str().cmp(other.as_str())
371    }
372}
373
374impl<A: RefCount> PartialOrd for ArcColdStringInner<A> {
375    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
376        Some(self.cmp(other))
377    }
378}
379
380impl<A: RefCount> str::FromStr for ArcColdStringInner<A> {
381    type Err = core::convert::Infallible;
382
383    fn from_str(s: &str) -> Result<Self, Self::Err> {
384        Ok(Self::new(s))
385    }
386}
387
388#[cfg(feature = "serde")]
389impl<A: RefCount> serde::Serialize for ArcColdStringInner<A> {
390    fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
391        serializer.serialize_str(self.as_str())
392    }
393}
394
395#[cfg(feature = "serde")]
396impl<'de, A: RefCount> serde::Deserialize<'de> for ArcColdStringInner<A> {
397    fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
398        let s = String::deserialize(deserializer)?;
399        Ok(Self::new(s))
400    }
401}
402
403unsafe impl<A: RefCount> Send for ArcColdStringInner<A> {}
404unsafe impl<A: RefCount> Sync for ArcColdStringInner<A> {}
405
406#[cfg(test)]
407mod tests {
408    use super::*;
409    use core::mem::{align_of, size_of};
410
411    macro_rules! each_ref_count {
412        ($test:ident) => {
413            $test::<AtomicU8>();
414            $test::<AtomicU16>();
415            $test::<AtomicU32>();
416            $test::<AtomicUsize>();
417        };
418    }
419
420    fn assert_layout<A: RefCount>() {
421        assert_eq!(size_of::<ArcColdStringInner<A>>(), size_of::<usize>());
422        assert_eq!(
423            size_of::<Option<ArcColdStringInner<A>>>(),
424            size_of::<ArcColdStringInner<A>>()
425        );
426        assert_eq!(size_of::<crate::heap::VintStringInner<A>>(), size_of::<A>());
427        assert_eq!(
428            align_of::<crate::heap::VintStringInner<A>>(),
429            align_of::<A>()
430        );
431    }
432
433    #[test]
434    fn layout() {
435        each_ref_count!(assert_layout);
436    }
437
438    fn assert_inline_and_heap_clone<A: RefCount>() {
439        let inline = ArcColdStringInner::<A>::new("tiny");
440        let inline_clone = inline.clone();
441        assert!(inline.is_inline());
442        assert_eq!(inline, inline_clone);
443
444        let heap = ArcColdStringInner::<A>::new("a string longer than one machine word");
445        let heap_clone = heap.clone();
446        assert!(!heap.is_inline());
447        assert_eq!(heap.encoded.addr(), heap_clone.encoded.addr());
448        assert_eq!(
449            heap.encoded.heap_ptr().as_ptr() as usize
450                % core::cmp::max(align_of::<A>(), crate::heap::HEAP_ALIGN),
451            0
452        );
453        assert_eq!(heap.count().refs(), 2);
454        drop(heap_clone);
455        assert_eq!(heap.count().refs(), 1);
456    }
457
458    #[test]
459    fn inline_and_heap_clone() {
460        each_ref_count!(assert_inline_and_heap_clone);
461    }
462
463    fn assert_clones_across_threads<A: RefCount>() {
464        let value = ArcColdStringInner::<A>::new("a shared string longer than one machine word");
465        let threads: alloc::vec::Vec<_> = (0..8)
466            .map(|_| {
467                let clone = value.clone();
468                std::thread::spawn(move || {
469                    assert_eq!(clone, "a shared string longer than one machine word")
470                })
471            })
472            .collect();
473
474        for thread in threads {
475            thread.join().unwrap();
476        }
477        assert_eq!(value.count().refs(), 1);
478    }
479
480    #[test]
481    fn clones_across_threads() {
482        each_ref_count!(assert_clones_across_threads);
483    }
484
485    #[cfg(all(loom, target_arch = "x86_64"))]
486    fn model_clone_drop<A: RefCount>() {
487        loom::model(|| {
488            const TEXT: &str = "a shared string longer than one machine word";
489
490            let value = ArcColdStringInner::<A>::new(TEXT);
491            let left = value.clone();
492            let right = value.clone();
493
494            let left = loom::thread::spawn(move || {
495                let clone = left.clone();
496                assert_eq!(clone.as_str(), TEXT);
497                drop(clone);
498                drop(left);
499            });
500            let right = loom::thread::spawn(move || {
501                let clone = right.clone();
502                assert_eq!(clone.as_str(), TEXT);
503                drop(right);
504                drop(clone);
505            });
506
507            left.join().unwrap();
508            right.join().unwrap();
509            assert_eq!(value.count().refs(), 1);
510        });
511    }
512
513    #[cfg(all(loom, target_arch = "x86_64"))]
514    fn model_final_drop<A: RefCount>() {
515        loom::model(|| {
516            let first =
517                ArcColdStringInner::<A>::new("a shared string longer than one machine word");
518            let second = first.clone();
519
520            let first = loom::thread::spawn(move || drop(first));
521            let second = loom::thread::spawn(move || drop(second));
522
523            first.join().unwrap();
524            second.join().unwrap();
525        });
526    }
527
528    #[cfg(all(loom, target_arch = "x86_64"))]
529    fn model_overflow_race<A: RefCount>() {
530        loom::model(|| {
531            let count = loom::sync::Arc::new(A::near_overflow());
532            let expected_refs = count.refs();
533            let increment = count.clone();
534            let decrement = count.clone();
535
536            let increment = loom::thread::spawn(move || increment.increment());
537            let decrement = loom::thread::spawn(move || assert!(!decrement.decrement()));
538
539            increment.join().unwrap();
540            decrement.join().unwrap();
541
542            if count.is_immortal() {
543                assert!(!count.decrement());
544                assert!(count.is_immortal());
545            } else {
546                assert_eq!(count.refs(), expected_refs);
547            }
548        });
549    }
550
551    #[cfg(all(loom, target_arch = "x86_64"))]
552    #[test]
553    fn loom_clone_drop() {
554        each_ref_count!(model_clone_drop);
555        each_ref_count!(model_final_drop);
556        each_ref_count!(model_overflow_race);
557    }
558
559    #[cfg(not(all(loom, target_arch = "x86_64")))]
560    fn assert_refcount_edges<A: RefCount>() {
561        let count = A::new();
562        assert_eq!(count.refs(), 1);
563        count.increment();
564        assert_eq!(count.refs(), 2);
565        assert!(!count.decrement());
566        assert_eq!(count.refs(), 1);
567        assert!(count.decrement());
568        assert_eq!(count.refs(), 0);
569
570        let count = A::near_overflow();
571        count.increment();
572        assert!(count.is_immortal());
573        assert!(!count.decrement());
574        assert!(count.is_immortal());
575
576        let count = A::min_immortal();
577        assert!(count.is_immortal());
578        assert!(!count.decrement());
579        assert!(count.is_immortal());
580    }
581
582    #[cfg(not(all(loom, target_arch = "x86_64")))]
583    #[test]
584    fn refcount_edges() {
585        each_ref_count!(assert_refcount_edges);
586    }
587
588    #[test]
589    fn const_inline_matches_runtime() {
590        macro_rules! assert_const_inline {
591            ($atomic:ty) => {{
592                const VALUE: ArcColdStringInner<$atomic> =
593                    ArcColdStringInner::<$atomic>::new_inline_const("cold");
594                assert_eq!(VALUE, ArcColdStringInner::<$atomic>::new("cold"));
595            }};
596        }
597
598        assert_const_inline!(AtomicU8);
599        assert_const_inline!(AtomicU16);
600        assert_const_inline!(AtomicU32);
601        assert_const_inline!(AtomicUsize);
602    }
603
604    #[cfg(feature = "serde")]
605    fn assert_serde_roundtrip<A: RefCount>() {
606        use serde_test::{assert_tokens, Token};
607
608        let value = ArcColdStringInner::<A>::new("a shared string longer than one machine word");
609        assert_tokens(
610            &value,
611            &[Token::Str("a shared string longer than one machine word")],
612        );
613    }
614
615    #[cfg(feature = "serde")]
616    #[test]
617    fn serde_roundtrip_shape() {
618        each_ref_count!(assert_serde_roundtrip);
619    }
620}