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
109pub type ArcColdString = ArcColdStringInner<AtomicUsize>;
124
125pub type ArcColdString8 = ArcColdStringInner<AtomicU8>;
130
131pub type ArcColdString16 = ArcColdStringInner<AtomicU16>;
136
137pub 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 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 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 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}