Skip to main content

vfs/
temp_clone.rs

1// Copyright 2023 The Fuchsia Authors. All rights reserved.
2// Use of this source code is governed by a BSD-style license that can be
3// found in the LICENSE file.
4
5use fuchsia_sync::{Condvar, Mutex};
6use std::cell::UnsafeCell;
7use std::collections::hash_map::Entry;
8use std::collections::{HashMap, VecDeque};
9use std::marker::PhantomData;
10use std::mem::ManuallyDrop;
11use std::ops::Deref;
12use std::sync::{Arc, OnceLock, Weak};
13
14use zx::sys::zx_handle_t;
15
16/// A wrapper around zircon handles that allows them to be temporarily cloned. These temporary
17/// clones can be used with `unblock` below which requires callbacks with static lifetime.  This is
18/// similar to Arc<T>, except that whilst there are no clones, there is no memory overhead, and
19/// there's no performance overhead to use them just as you would without the wrapper, except for a
20/// small overhead when they are dropped. The wrapper ensures that the handle is only dropped when
21/// there are no references.
22pub struct TempClonable<T: fidl::AsHandleRef>(ManuallyDrop<T>);
23
24impl<T: fidl::AsHandleRef> TempClonable<T> {
25    /// Returns a new handle that can be temporarily cloned.
26    pub fn new(handle: T) -> Self {
27        Self(ManuallyDrop::new(handle))
28    }
29}
30
31impl<T: fidl::AsHandleRef> Deref for TempClonable<T> {
32    type Target = T;
33
34    fn deref(&self) -> &T {
35        &self.0
36    }
37}
38
39impl<T: fidl::AsHandleRef> TempClonable<T> {
40    /// Creates a temporary clone of the handle. The clone should only exist temporarily.
41    ///
42    /// # Panics
43    ///
44    /// Panics if the handle is invalid.
45    pub fn temp_clone(&self) -> TempClone<T> {
46        assert!(!self.as_handle_ref().is_invalid());
47        let mut clones = clones().lock();
48        let raw_handle = self.0.as_handle_ref().raw_handle();
49        TempClone {
50            handle: match clones.entry(raw_handle) {
51                Entry::Occupied(mut o) => {
52                    if let Some(clone) = o.get().upgrade() {
53                        clone
54                    } else {
55                        // The last strong reference was dropped but the entry hasn't been removed
56                        // yet. This must be racing with `TempHandle::drop`. Replace the
57                        // `TempHandle`.
58                        let clone =
59                            Arc::new(TempHandle { raw_handle, tombstone: UnsafeCell::new(false) });
60                        *o.get_mut() = Arc::downgrade(&clone);
61                        clone
62                    }
63                }
64                Entry::Vacant(v) => {
65                    let clone =
66                        Arc::new(TempHandle { raw_handle, tombstone: UnsafeCell::new(false) });
67                    v.insert(Arc::downgrade(&clone));
68                    clone
69                }
70            },
71            marker: PhantomData,
72        }
73    }
74}
75
76impl<T: fidl::AsHandleRef> Drop for TempClonable<T> {
77    fn drop(&mut self) {
78        if let Some(handle) =
79            clones().lock().remove(&self.0.as_handle_ref().raw_handle()).and_then(|c| c.upgrade())
80        {
81            // There are still some temporary clones alive, so mark the handle with a tombstone.
82
83            // SAFETY: This is the only unsafe place where we access `tombstone`. We're are holding
84            // the clones lock which ensures no other thread is concurrently accessing it, but it
85            // wouldn't normally happen anyway because it would mean there were multiple
86            // TempClonable instances wrapping the same handle, which shouldn't happen.
87            unsafe { *handle.tombstone.get() = true };
88            return;
89        }
90
91        // SAFETY: There are no temporary clones, so we can drop the handle now. No more clones can
92        // be made and it should be clear we meet the safety requirements of ManuallyDrop.
93        unsafe { ManuallyDrop::drop(&mut self.0) }
94    }
95}
96
97type Clones = Mutex<HashMap<zx_handle_t, Weak<TempHandle>>>;
98
99/// Returns the global instance which keeps track of temporary clones.
100fn clones() -> &'static Clones {
101    static CLONES: OnceLock<Clones> = OnceLock::new();
102    CLONES.get_or_init(|| Mutex::new(HashMap::new()))
103}
104
105pub struct TempClone<T> {
106    handle: Arc<TempHandle>,
107    marker: PhantomData<T>,
108}
109
110impl<T> Deref for TempClone<T> {
111    type Target = T;
112
113    fn deref(&self) -> &T {
114        // SAFETY: T is repr(transparent) and stores zx_handle_t.
115        unsafe { std::mem::transmute::<&zx_handle_t, &T>(&self.handle.raw_handle) }
116    }
117}
118
119struct TempHandle {
120    raw_handle: zx_handle_t,
121    tombstone: UnsafeCell<bool>,
122}
123
124unsafe impl Send for TempHandle {}
125unsafe impl Sync for TempHandle {}
126
127impl Drop for TempHandle {
128    fn drop(&mut self) {
129        if *self.tombstone.get_mut() {
130            // SAFETY: The primary handle has been dropped and it is our job to clean up the
131            // handle. There are no memory safety issues here.
132            unsafe { fidl::NullableHandle::from_raw(self.raw_handle) };
133        } else {
134            if let Entry::Occupied(o) = clones().lock().entry(self.raw_handle) {
135                // There's a small window where another TempHandle could have been inserted, so
136                // before removing this entry, check for a match.
137                if std::ptr::eq(o.get().as_ptr(), self) {
138                    o.remove_entry();
139                }
140            }
141        }
142    }
143}
144
145/// This is similar to fuchsia-async's unblock except that it used a fixed size thread pool which
146/// has the advantage of not making traces difficult to decipher because of many threads being
147/// spawned.
148pub async fn unblock<T: 'static + Send>(f: impl FnOnce() -> T + Send + 'static) -> T {
149    const NUM_THREADS: u8 = 2;
150
151    struct State {
152        queue: Mutex<VecDeque<Box<dyn FnOnce() + Send + 'static>>>,
153        cvar: Condvar,
154    }
155
156    static STATE: OnceLock<State> = OnceLock::new();
157
158    let mut start_threads = false;
159    let state = STATE.get_or_init(|| {
160        start_threads = true;
161        State { queue: Mutex::new(VecDeque::new()), cvar: Condvar::new() }
162    });
163
164    if start_threads {
165        for _ in 0..NUM_THREADS {
166            std::thread::spawn(|| {
167                loop {
168                    let item = {
169                        let mut queue = state.queue.lock();
170                        loop {
171                            if let Some(item) = queue.pop_front() {
172                                break item;
173                            }
174                            state.cvar.wait(&mut queue);
175                        }
176                    };
177                    item();
178                }
179            });
180        }
181    }
182
183    let (tx, rx) = futures::channel::oneshot::channel();
184    state.queue.lock().push_back(Box::new(move || {
185        let _ = tx.send(f());
186    }));
187    state.cvar.notify_one();
188
189    rx.await.unwrap()
190}
191
192#[cfg(target_os = "fuchsia")]
193#[cfg(test)]
194mod tests {
195    use super::{TempClonable, clones};
196
197    use std::sync::Arc;
198
199    #[test]
200    fn test_temp_clone() {
201        let parent_vmo = zx::Vmo::create(100).expect("create failed");
202
203        {
204            let temp_clone = {
205                let vmo = TempClonable::new(
206                    parent_vmo
207                        .create_child(zx::VmoChildOptions::REFERENCE, 0, 0)
208                        .expect("create_child failed"),
209                );
210
211                vmo.write(b"foo", 0).expect("write failed");
212                {
213                    // Create and read from a temporary clone.
214                    let temp_clone2 = vmo.temp_clone();
215                    assert_eq!(
216                        &temp_clone2.read_to_vec::<u8>(0, 3).expect("read_to_vec failed"),
217                        b"foo"
218                    );
219                }
220
221                // We should still be able to read from the primary handle.
222                assert_eq!(&vmo.read_to_vec::<u8>(0, 3).expect("read_to_vec failed"), b"foo");
223
224                // Create another vmo which should get cleaned up when the primary handle is
225                // dropped.
226                let vmo2 = TempClonable::new(
227                    parent_vmo
228                        .create_child(zx::VmoChildOptions::REFERENCE, 0, 0)
229                        .expect("create_child failed"),
230                );
231                // Create and immediately drop a temporary clone.
232                vmo2.temp_clone();
233
234                // Take another clone that will get dropped after we take the clone below.
235                let _clone1 = vmo.temp_clone();
236
237                // And return another clone.
238                vmo.temp_clone()
239            };
240
241            // The primary handle has been dropped, but we should still be able to
242            // read via temp_clone.
243            assert_eq!(&temp_clone.read_to_vec::<u8>(0, 3).expect("read_to_vec failed"), b"foo");
244        }
245
246        // Make sure that all the VMOs got properly cleaned up.
247        parent_vmo
248            .wait_one(zx::Signals::VMO_ZERO_CHILDREN, zx::MonotonicInstant::INFINITE)
249            .expect("wait for zero children failed");
250        assert_eq!(parent_vmo.info().expect("info failed").num_children, 0);
251        assert!(clones().lock().is_empty());
252    }
253
254    #[test]
255    fn test_race() {
256        let parent_vmo = zx::Vmo::create(100).expect("create failed");
257
258        {
259            let vmo = Arc::new(TempClonable::new(
260                parent_vmo
261                    .create_child(zx::VmoChildOptions::REFERENCE, 0, 0)
262                    .expect("create_child failed"),
263            ));
264            vmo.write(b"foo", 0).expect("write failed");
265
266            let vmo_clone = vmo.clone();
267
268            let t1 = std::thread::spawn(move || {
269                for _ in 0..1000 {
270                    assert_eq!(
271                        &vmo.temp_clone().read_to_vec::<u8>(0, 3).expect("read_to_vec failed"),
272                        b"foo"
273                    );
274                }
275            });
276
277            let t2 = std::thread::spawn(move || {
278                for _ in 0..1000 {
279                    assert_eq!(
280                        &vmo_clone
281                            .temp_clone()
282                            .read_to_vec::<u8>(0, 3)
283                            .expect("read_to_vec failed"),
284                        b"foo"
285                    );
286                }
287            });
288
289            let _ = t1.join();
290            let _ = t2.join();
291        }
292
293        // Make sure that all the VMOs got properly cleaned up.
294        parent_vmo
295            .wait_one(zx::Signals::VMO_ZERO_CHILDREN, zx::MonotonicInstant::INFINITE)
296            .expect("wait for zero children failed");
297        assert_eq!(parent_vmo.info().expect("info failed").num_children, 0);
298        assert!(clones().lock().is_empty());
299    }
300}