Skip to main content

nest_guard/
lib.rs

1//! This crate works around the following weakness in standard Rust. It is
2//! generally impossible to write:
3//!
4//!   // Say that you're working with a Weak<RefCell<u32>> for safe, shared
5//!   // mutability.
6//!   let x = Rc::new(RefCell::new(0));
7//!   let x = Rc::downgrade(&x);
8//!
9//!   // This doesn't work. The Rc returned by x.upgrade() is dropped at the end
10//!   // of the statement, and therefore ref_inner can't be used.
11//!   let ref_inner = x.upgrade().unwrap().borrow();
12//!   assert_eq!(0, *ref_inner);
13//!
14//! With this crate, replacing `borrow` with `nest_borrow` Just Works
15//!
16//!   let x = Rc::new(RefCell::new(0));
17//!   let x = Rc::downgrade(&x);
18//!
19//!   let ref_inner = x.upgrade().unwrap().nest_borrow();
20//!   assert_eq!(0, *ref_inner);
21//!
22//! The implementation relies on one `unsafe` function to prevent rustc from
23//! being overly strict.
24//!
25//! The implementation guarantees that the "outer temporary" (Rc in the above
26//! example) does in fact outlive the reference to it.
27//!
28//! The implementation has no runtime overhead.
29//!
30//! Arbitrary nesting depth is supported
31//!
32//!   let x = RefCell::new(RefCell::new(RefCell::new(0)));
33//!   let ref_inner = x.borrow().nest_borrow().nest_borrow();
34//!   assert_eq!(0, *ref_inner);
35//!
36//! `nest_*` functions also work fine in the initial position of the chain, in
37//! which position they behave like their unprefixed analogues. The following
38//! two lines are identical.
39//!
40//!   let ref_inner = x.     borrow().nest_borrow().nest_borrow();
41//!   let ref_inner = x.nest_borrow().nest_borrow().nest_borrow();
42//!
43//! If you ever have had to work with the miserable experience of reassigning a
44//! stack of guards, you may recognize that this code doesn't compile
45//!
46//!   let x = RefCell::new(RefCell::new(0));
47//!   let ys: Vec<_> = (1..10).map(|i| RefCell::new(RefCell::new(i))).collect();
48//!
49//!   let mut ref1 = x.borrow();
50//!   let mut ref2 = ref1.borrow_mut();
51//!
52//!   for y in ys.iter() {
53//!     let mut yref1 = y.borrow();
54//!     let mut yref2 = yref1.borrow_mut();
55//!     mem::swap(ref2.deref_mut(), yref2.deref_mut());
56//!     ref2 = yref2;
57//!     ref1 = yref1;
58//!   }
59//!
60//! However, the analogue using this crate again Just Works
61//!
62//!   let x = RefCell::new(RefCell::new(0));
63//!   let ys: Vec<_> = (1..10).map(|i| RefCell::new(RefCell::new(i))).collect();
64//!
65//!   let mut ref1 = x.borrow().nest_borrow_mut();
66//!
67//!   for y in ys.iter() {
68//!     let mut yref = y.borrow().nest_borrow_mut();
69//!     mem::swap(ref1.deref_mut(), yref.deref_mut());
70//!     ref1 = yref;
71//!   }
72//!
73//! The following methods are provided, via extension traits
74//!   std::cell::RefCell::nest_borrow
75//!   std::cell::RefCell::nest_borrow_mut
76//!   std::cell::RefCell::try_nest_borrow
77//!   std::cell::RefCell::try_nest_borrow_mut
78//!   std::rc::Weak::nest_upgrade
79//!   std::sync::Weak::nest_upgrade
80//!   std::sync::Mutex::nest_lock
81//!   std::sync::Mutex::nest_try_lock
82//!   std::sync::RwLock::nest_try_read
83//!   std::sync::RwLock::nest_try_write
84
85use std::ops::{Deref, DerefMut};
86
87pub use self::cell::*;
88pub use self::rc::*;
89pub use self::sync::*;
90
91/// NOTE that drop order is guaranteed first-to-last, so the inner ref is
92/// dropped before the outer ref, which we rely on in our unsafe transmutes.
93pub struct Nested<T, Inner: Deref<Target = T>, Outer> {
94    inner: Inner,
95    #[allow(dead_code)]
96    outer: Outer,
97}
98impl<T, Inner: Deref<Target = T>, Outer> Deref for Nested<T, Inner, Outer> {
99    type Target = T;
100    fn deref(&self) -> &Self::Target {
101        self.inner.deref()
102    }
103}
104impl<T, Inner: DerefMut<Target = T>, Outer> DerefMut for Nested<T, Inner, Outer> {
105    fn deref_mut(&mut self) -> &mut Self::Target {
106        self.inner.deref_mut()
107    }
108}
109
110/// # Safety
111///   At each usage site in this file, the justification is the same. We work
112///   around rustc's inability to understand self-referential structs by casting
113///   away the lifetime, and manually ensure that the referenced object outlives
114///   the reference by bundling them both into the `Nested` struct.
115unsafe fn remove_lifetime<'b, T>(x: &T) -> &'b T {
116    unsafe { &*(x as *const T) }
117}
118
119mod cell {
120    use std::cell::*;
121
122    use super::*;
123
124    pub trait NestedRefCell<'a, T>: Deref<Target = RefCell<T>> + Sized + 'a {
125        fn nest_borrow(self) -> Nested<T, Ref<'a, T>, Self> {
126            let me = unsafe { remove_lifetime(&self) };
127            let inner = RefCell::borrow(me);
128            Nested { inner, outer: self }
129        }
130        fn nest_try_borrow(self) -> Result<Nested<T, Ref<'a, T>, Self>, BorrowError> {
131            let me = unsafe { remove_lifetime(&self) };
132            let inner = RefCell::try_borrow(me)?;
133            Ok(Nested { inner, outer: self })
134        }
135        fn nest_borrow_mut(self) -> Nested<T, RefMut<'a, T>, Self> {
136            let me = unsafe { remove_lifetime(&self) };
137            let inner = RefCell::borrow_mut(me);
138            Nested { inner, outer: self }
139        }
140        fn nest_try_borrow_mut(self) -> Result<Nested<T, RefMut<'a, T>, Self>, BorrowMutError> {
141            let me = unsafe { remove_lifetime(&self) };
142            let inner = RefCell::try_borrow_mut(me)?;
143            Ok(Nested { inner, outer: self })
144        }
145    }
146    impl<'a, T, Outer: Deref<Target = RefCell<T>> + Sized + 'a> NestedRefCell<'a, T> for Outer {}
147
148    #[cfg(test)]
149    mod test {
150        use super::*;
151
152        #[test]
153        fn test_many_refcell() {
154            let x = RefCell::new(RefCell::new(RefCell::new(0)));
155            {
156                let z = x.borrow().nest_borrow().nest_borrow();
157                assert_eq!(0, *z);
158            }
159            {
160                let z = x
161                    .try_borrow()
162                    .unwrap()
163                    .nest_try_borrow()
164                    .unwrap()
165                    .nest_try_borrow()
166                    .unwrap();
167                assert_eq!(0, *z);
168            }
169            {
170                let mut z = x.borrow().nest_borrow().nest_borrow_mut();
171                *z = 1;
172                assert_eq!(1, *z);
173            }
174            {
175                {
176                    let mut y = x.borrow().nest_borrow_mut();
177                    *y = RefCell::new(2);
178                }
179                let z = x.borrow().nest_borrow().nest_borrow();
180                assert_eq!(2, *z);
181            }
182        }
183    }
184}
185mod rc {
186    use std::rc::*;
187
188    use super::*;
189
190    pub trait NestedRcWeak<'a, T>: Deref<Target = Weak<T>> + Sized + 'a {
191        fn nest_upgrade(self) -> Option<Nested<T, Rc<T>, Self>> {
192            let inner = Weak::upgrade(&self)?;
193            Some(Nested { inner, outer: self })
194        }
195    }
196    impl<'a, T, Outer: Deref<Target = Weak<T>> + Sized + 'a> NestedRcWeak<'a, T> for Outer {}
197    #[cfg(test)]
198    mod test {
199        use super::*;
200
201        #[test]
202        fn test_many_rc() {
203            let x1: Rc<i32> = Rc::new(0);
204            let y1: Weak<i32> = Rc::downgrade(&x1);
205
206            let x2: Rc<Weak<i32>> = Rc::new(y1);
207            let y2: Weak<Weak<i32>> = Rc::downgrade(&x2);
208
209            let x3: Rc<Weak<Weak<i32>>> = Rc::new(y2);
210            let y3: Weak<Weak<Weak<i32>>> = Rc::downgrade(&x3);
211
212            let z = y3
213                .upgrade()
214                .unwrap()
215                .nest_upgrade()
216                .unwrap()
217                .nest_upgrade()
218                .unwrap();
219            assert_eq!(0, *z);
220        }
221    }
222}
223
224mod sync {
225    use std::sync::*;
226
227    use super::*;
228
229    pub trait NestedArcWeak<'a, T>: Deref<Target = Weak<T>> + Sized + 'a {
230        fn nest_upgrade(self) -> Option<Nested<T, Arc<T>, Self>> {
231            let inner = Weak::upgrade(&self)?;
232            Some(Nested { inner, outer: self })
233        }
234    }
235    impl<'a, T, Outer: Deref<Target = Weak<T>> + Sized + 'a> NestedArcWeak<'a, T> for Outer {}
236
237    pub trait NestedMutex<'a, T>: Deref<Target = Mutex<T>> + Sized + 'a {
238        fn nest_lock(self) -> LockResult<Nested<T, MutexGuard<'a, T>, Self>> {
239            let me = unsafe { remove_lifetime(&self) };
240            match me.lock() {
241                Ok(inner) => Ok(Nested { inner, outer: self }),
242                Err(err) => {
243                    let inner = err.into_inner();
244                    Err(PoisonError::new(Nested { inner, outer: self }))
245                }
246            }
247        }
248        fn nest_try_lock(self) -> TryLockResult<Nested<T, MutexGuard<'a, T>, Self>> {
249            let me = unsafe { remove_lifetime(&self) };
250            match me.try_lock() {
251                Ok(inner) => Ok(Nested { inner, outer: self }),
252                Err(err) => match err {
253                    TryLockError::Poisoned(err) => {
254                        let inner = err.into_inner();
255                        Err(TryLockError::Poisoned(PoisonError::new(Nested {
256                            inner,
257                            outer: self,
258                        })))
259                    }
260                    TryLockError::WouldBlock => Err(TryLockError::WouldBlock),
261                },
262            }
263        }
264    }
265    impl<'a, T, Outer: Deref<Target = Mutex<T>> + Sized + 'a> NestedMutex<'a, T> for Outer {}
266
267    pub trait NestedRwLock<'a, T>: Deref<Target = RwLock<T>> + Sized + 'a {
268        fn nest_try_read(self) -> TryLockResult<Nested<T, RwLockReadGuard<'a, T>, Self>> {
269            let me = unsafe { remove_lifetime(&self) };
270            match me.try_read() {
271                Ok(inner) => Ok(Nested { inner, outer: self }),
272                Err(err) => match err {
273                    TryLockError::Poisoned(err) => {
274                        let inner = err.into_inner();
275                        Err(TryLockError::Poisoned(PoisonError::new(Nested {
276                            inner,
277                            outer: self,
278                        })))
279                    }
280                    TryLockError::WouldBlock => Err(TryLockError::WouldBlock),
281                },
282            }
283        }
284        fn nest_try_write(self) -> TryLockResult<Nested<T, RwLockWriteGuard<'a, T>, Self>> {
285            let me = unsafe { remove_lifetime(&self) };
286            match me.try_write() {
287                Ok(inner) => Ok(Nested { inner, outer: self }),
288                Err(err) => match err {
289                    TryLockError::Poisoned(err) => {
290                        let inner = err.into_inner();
291                        Err(TryLockError::Poisoned(PoisonError::new(Nested {
292                            inner,
293                            outer: self,
294                        })))
295                    }
296                    TryLockError::WouldBlock => Err(TryLockError::WouldBlock),
297                },
298            }
299        }
300    }
301    impl<'a, T, Outer: Deref<Target = RwLock<T>> + Sized + 'a> NestedRwLock<'a, T> for Outer {}
302
303    #[cfg(test)]
304    mod test {
305        use super::*;
306
307        #[test]
308        fn test_many_arc() {
309            let x1: Arc<i32> = Arc::new(0);
310            let y1: Weak<i32> = Arc::downgrade(&x1);
311
312            let x2: Arc<Weak<i32>> = Arc::new(y1);
313            let y2: Weak<Weak<i32>> = Arc::downgrade(&x2);
314
315            let x3: Arc<Weak<Weak<i32>>> = Arc::new(y2);
316            let y3: Weak<Weak<Weak<i32>>> = Arc::downgrade(&x3);
317
318            let z = y3
319                .upgrade()
320                .unwrap()
321                .nest_upgrade()
322                .unwrap()
323                .nest_upgrade()
324                .unwrap();
325            assert_eq!(0, *z);
326        }
327
328        #[test]
329        fn test_many_rwlock() {
330            let x = RwLock::new(RwLock::new(RwLock::new(0)));
331            {
332                let z = x
333                    .try_read()
334                    .unwrap()
335                    .nest_try_read()
336                    .unwrap()
337                    .nest_try_read()
338                    .unwrap();
339                assert_eq!(0, *z);
340            }
341            {
342                let z = x
343                    .try_read()
344                    .unwrap()
345                    .nest_try_read()
346                    .unwrap()
347                    .nest_try_read()
348                    .unwrap();
349                assert_eq!(0, *z);
350            }
351            {
352                let mut z = x
353                    .try_read()
354                    .unwrap()
355                    .nest_try_read()
356                    .unwrap()
357                    .nest_try_write()
358                    .unwrap();
359                *z = 1;
360                assert_eq!(1, *z);
361            }
362            {
363                {
364                    let mut y = x.try_read().unwrap().nest_try_write().unwrap();
365                    *y = RwLock::new(2);
366                }
367                let z = x
368                    .try_read()
369                    .unwrap()
370                    .nest_try_read()
371                    .unwrap()
372                    .nest_try_read()
373                    .unwrap();
374                assert_eq!(2, *z);
375            }
376        }
377        #[test]
378        fn test_many_mutex() {
379            let x = RwLock::new(RwLock::new(RwLock::new(0)));
380            {
381                let z = x
382                    .try_read()
383                    .unwrap()
384                    .nest_try_read()
385                    .unwrap()
386                    .nest_try_read()
387                    .unwrap();
388                assert_eq!(0, *z);
389            }
390            {
391                let z = x
392                    .try_read()
393                    .unwrap()
394                    .nest_try_read()
395                    .unwrap()
396                    .nest_try_read()
397                    .unwrap();
398                assert_eq!(0, *z);
399            }
400            {
401                let mut z = x
402                    .try_read()
403                    .unwrap()
404                    .nest_try_read()
405                    .unwrap()
406                    .nest_try_write()
407                    .unwrap();
408                *z = 1;
409                assert_eq!(1, *z);
410            }
411            {
412                {
413                    let mut y = x.try_read().unwrap().nest_try_write().unwrap();
414                    *y = RwLock::new(2);
415                }
416                let z = x
417                    .try_read()
418                    .unwrap()
419                    .nest_try_read()
420                    .unwrap()
421                    .nest_try_read()
422                    .unwrap();
423                assert_eq!(2, *z);
424            }
425        }
426    }
427}
428
429#[cfg(test)]
430mod test {
431    use std::cell::*;
432    use std::mem;
433    use std::rc::*;
434
435    use super::*;
436
437    #[test]
438    fn test_rc_refcell() {
439        let x = Rc::new(RefCell::new(0));
440        let x = Rc::downgrade(&x);
441        assert_eq!(0, *x.upgrade().unwrap().borrow());
442        {
443            let z = x.upgrade().unwrap().nest_borrow();
444            assert_eq!(0, *z);
445        }
446    }
447
448    // #[test]
449    // fn test_reassign_refcell_stack_does_not_compile() {
450    //   let x = RefCell::new(RefCell::new(0));
451    //   let ys: Vec<_> = (1..10).map(|i|
452    // RefCell::new(RefCell::new(i))).collect();
453
454    //   let mut ref1 = x.borrow();
455    //   let mut ref2 = ref1.borrow_mut();
456
457    //   for y in ys.iter() {
458    //     let mut yref1 = y.borrow();
459    //     let mut yref2 = yref1.borrow_mut();
460    //     mem::swap(ref2.deref_mut(), yref2.deref_mut());
461    //     ref2 = yref2;
462    //     ref1 = yref1;
463    //   }
464    // }
465
466    #[test]
467    fn test_reassign_refcell_stack() {
468        let x = RefCell::new(RefCell::new(0));
469        let ys: Vec<_> = (1..10).map(|i| RefCell::new(RefCell::new(i))).collect();
470
471        let mut ref1 = x.borrow().nest_borrow_mut();
472
473        for y in ys.iter() {
474            let mut yref = y.borrow().nest_borrow_mut();
475            mem::swap(ref1.deref_mut(), yref.deref_mut());
476            ref1 = yref;
477        }
478    }
479}