Skip to main content

gpui/app/
entity_map.rs

1use crate::{App, AppContext, GpuiBorrow, VisualContext, Window, seal::Sealed};
2use anyhow::{Context as _, Result};
3use collections::FxHashSet;
4use derive_more::{Deref, DerefMut};
5use parking_lot::{RwLock, RwLockUpgradableReadGuard};
6use slotmap::{KeyData, SecondaryMap, SlotMap};
7use std::{
8    any::{Any, TypeId, type_name},
9    cell::RefCell,
10    cmp::Ordering,
11    fmt::{self, Display},
12    hash::{Hash, Hasher},
13    marker::PhantomData,
14    num::NonZeroU64,
15    sync::{
16        Arc, Weak,
17        atomic::{AtomicU64, AtomicUsize, Ordering::SeqCst},
18    },
19    thread::panicking,
20};
21
22use super::Context;
23use crate::util::atomic_incr_if_not_zero;
24#[cfg(any(test, gpui_leak_detection))]
25use collections::HashMap;
26
27slotmap::new_key_type! {
28    /// A unique identifier for a entity across the application.
29    pub struct EntityId;
30}
31
32impl From<u64> for EntityId {
33    fn from(value: u64) -> Self {
34        Self(KeyData::from_ffi(value))
35    }
36}
37
38impl EntityId {
39    /// Converts this entity id to a [NonZeroU64]
40    pub fn as_non_zero_u64(self) -> NonZeroU64 {
41        NonZeroU64::new(self.0.as_ffi()).unwrap()
42    }
43
44    /// Converts this entity id to a [u64]
45    pub fn as_u64(self) -> u64 {
46        self.0.as_ffi()
47    }
48}
49
50impl Display for EntityId {
51    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
52        write!(f, "{}", self.as_u64())
53    }
54}
55
56pub(crate) struct EntityMap {
57    entities: SecondaryMap<EntityId, Box<dyn Any>>,
58    pub accessed_entities: RefCell<FxHashSet<EntityId>>,
59    ref_counts: Arc<RwLock<EntityRefCounts>>,
60}
61
62#[doc(hidden)]
63pub(crate) struct EntityRefCounts {
64    counts: SlotMap<EntityId, AtomicUsize>,
65    dropped_entity_ids: Vec<EntityId>,
66    #[cfg(any(test, gpui_leak_detection))]
67    leak_detector: LeakDetector,
68}
69
70pub(super) struct LeaseInner {
71    pub(super) entity: Option<Box<dyn Any>>,
72}
73
74impl EntityMap {
75    pub fn new() -> Self {
76        Self {
77            entities: SecondaryMap::new(),
78            accessed_entities: RefCell::new(FxHashSet::default()),
79            ref_counts: Arc::new(RwLock::new(EntityRefCounts {
80                counts: SlotMap::with_key(),
81                dropped_entity_ids: Vec::new(),
82                #[cfg(any(test, gpui_leak_detection))]
83                leak_detector: LeakDetector {
84                    next_handle_id: 0,
85                    entity_handles: HashMap::default(),
86                },
87            })),
88        }
89    }
90
91    #[doc(hidden)]
92    pub fn ref_counts_drop_handle(&self) -> Arc<RwLock<EntityRefCounts>> {
93        self.ref_counts.clone()
94    }
95
96    /// Captures a snapshot of all entities that currently have alive handles.
97    ///
98    /// The returned [`LeakDetectorSnapshot`] can later be passed to
99    /// [`assert_no_new_leaks`](Self::assert_no_new_leaks) to verify that no
100    /// entities created after the snapshot are still alive.
101    #[cfg(any(test, gpui_leak_detection))]
102    pub fn leak_detector_snapshot(&self) -> LeakDetectorSnapshot {
103        self.ref_counts.read().leak_detector.snapshot()
104    }
105
106    /// Asserts that no entities created after `snapshot` still have alive handles.
107    ///
108    /// See [`LeakDetector::assert_no_new_leaks`] for details.
109    #[cfg(any(test, gpui_leak_detection))]
110    pub fn assert_no_new_leaks(&self, snapshot: &LeakDetectorSnapshot) {
111        self.ref_counts
112            .read()
113            .leak_detector
114            .assert_no_new_leaks(snapshot)
115    }
116
117    /// Reserve a slot for an entity, which you can subsequently use with `insert`.
118    pub fn reserve<T: 'static>(&self) -> Slot<T> {
119        let id = self.ref_counts.write().counts.insert(1.into());
120        Slot(Entity::new(id, Arc::downgrade(&self.ref_counts)))
121    }
122
123    /// Insert an entity into a slot obtained by calling `reserve`.
124    pub fn insert<T>(&mut self, slot: Slot<T>, entity: T) -> Entity<T>
125    where
126        T: 'static,
127    {
128        let mut accessed_entities = self.accessed_entities.get_mut();
129        accessed_entities.insert(slot.entity_id);
130
131        let handle = slot.0;
132        self.entities.insert(handle.entity_id, Box::new(entity));
133        handle
134    }
135
136    /// Move an entity to the stack.
137    #[track_caller]
138    pub fn lease<T>(&mut self, pointer: &Entity<T>) -> Lease<T> {
139        Lease {
140            inner: self.lease_erased(pointer, type_name::<T>()),
141            id: pointer.entity_id,
142            entity_type: PhantomData,
143        }
144    }
145
146    /// Returns an entity after moving it to the stack.
147    pub fn end_lease<T>(&mut self, lease: Lease<T>) {
148        self.end_lease_erased(lease.id, lease.inner);
149    }
150
151    #[inline(always)]
152    pub fn read<T: 'static>(&self, entity: &Entity<T>) -> &T {
153        self.assert_valid_context(entity);
154        self.read_inner(entity.entity_id)
155            .and_then(|entity| entity.downcast_ref())
156            .unwrap_or_else(|| double_lease_panic("read", type_name::<T>()))
157    }
158
159    #[track_caller]
160    pub(super) fn lease_erased(&mut self, pointer: &AnyEntity, entity_type: &str) -> LeaseInner {
161        self.assert_valid_context(pointer);
162        let entity = Some(
163            self.lease_inner(pointer.entity_id)
164                .unwrap_or_else(|| double_lease_panic("update", entity_type)),
165        );
166        LeaseInner { entity }
167    }
168
169    pub(super) fn end_lease_erased(&mut self, entity_id: EntityId, mut lease: LeaseInner) {
170        self.end_lease_inner(entity_id, lease.entity.take().unwrap());
171    }
172
173    fn assert_valid_context(&self, entity: &AnyEntity) {
174        debug_assert!(
175            Weak::ptr_eq(&entity.entity_map, &Arc::downgrade(&self.ref_counts)),
176            "used a entity with the wrong context"
177        );
178    }
179
180    pub fn extend_accessed(&mut self, entities: &FxHashSet<EntityId>) {
181        self.accessed_entities
182            .get_mut()
183            .extend(entities.iter().copied());
184    }
185
186    pub fn clear_accessed(&mut self) {
187        self.accessed_entities.get_mut().clear();
188    }
189
190    pub fn take_dropped(&mut self) -> Vec<(EntityId, Box<dyn Any>)> {
191        let mut ref_counts = &mut *self.ref_counts.write();
192        let dropped_entity_ids = ref_counts.dropped_entity_ids.drain(..);
193        let mut accessed_entities = self.accessed_entities.get_mut();
194
195        dropped_entity_ids
196            .filter_map(|entity_id| {
197                let count = ref_counts.counts.remove(entity_id).unwrap();
198                debug_assert_eq!(
199                    count.load(SeqCst),
200                    0,
201                    "dropped an entity that was referenced"
202                );
203                accessed_entities.remove(&entity_id);
204                // If the EntityId was allocated with `Context::reserve`,
205                // the entity may not have been inserted.
206                Some((entity_id, self.entities.remove(entity_id)?))
207            })
208            .collect()
209    }
210
211    #[inline(never)]
212    fn read_inner(&self, entity_id: EntityId) -> Option<&dyn Any> {
213        let mut accessed_entities = self.accessed_entities.borrow_mut();
214        accessed_entities.insert(entity_id);
215        self.entities.get(entity_id).map(Box::as_ref)
216    }
217
218    #[inline(never)]
219    fn lease_inner(&mut self, entity_id: EntityId) -> Option<Box<dyn Any>> {
220        self.accessed_entities.get_mut().insert(entity_id);
221        self.entities.remove(entity_id)
222    }
223
224    #[inline(never)]
225    fn end_lease_inner(&mut self, entity_id: EntityId, entity: Box<dyn Any>) {
226        self.entities.insert(entity_id, entity);
227    }
228}
229
230#[track_caller]
231fn double_lease_panic(operation: &str, entity_type: &str) -> ! {
232    panic!("cannot {operation} {entity_type} while it is already being updated")
233}
234
235pub(crate) struct Lease<T> {
236    pub id: EntityId,
237    inner: LeaseInner,
238    entity_type: PhantomData<T>,
239}
240
241impl<T: 'static> core::ops::Deref for Lease<T> {
242    type Target = T;
243
244    fn deref(&self) -> &Self::Target {
245        self.inner.entity.as_ref().unwrap().downcast_ref().unwrap()
246    }
247}
248
249impl<T: 'static> core::ops::DerefMut for Lease<T> {
250    fn deref_mut(&mut self) -> &mut Self::Target {
251        self.inner.entity.as_mut().unwrap().downcast_mut().unwrap()
252    }
253}
254
255impl Drop for LeaseInner {
256    fn drop(&mut self) {
257        if self.entity.is_some() && !panicking() {
258            panic!("Leases must be ended with EntityMap::end_lease")
259        }
260    }
261}
262
263#[derive(Deref, DerefMut)]
264pub(crate) struct Slot<T>(Entity<T>);
265
266/// A dynamically typed reference to a entity, which can be downcast into a `Entity<T>`.
267pub struct AnyEntity {
268    pub(crate) entity_id: EntityId,
269    pub(crate) entity_type: TypeId,
270    entity_map: Weak<RwLock<EntityRefCounts>>,
271    #[cfg(any(test, gpui_leak_detection))]
272    handle_id: HandleId,
273}
274
275impl AnyEntity {
276    fn new(
277        id: EntityId,
278        entity_type: TypeId,
279        entity_map: Weak<RwLock<EntityRefCounts>>,
280        #[cfg(any(test, gpui_leak_detection))] type_name: &'static str,
281    ) -> Self {
282        Self {
283            entity_id: id,
284            entity_type,
285            #[cfg(any(test, gpui_leak_detection))]
286            handle_id: entity_map
287                .clone()
288                .upgrade()
289                .unwrap()
290                .write()
291                .leak_detector
292                .handle_created(id, Some(type_name)),
293            entity_map,
294        }
295    }
296
297    /// Returns the id associated with this entity.
298    #[inline]
299    pub fn entity_id(&self) -> EntityId {
300        self.entity_id
301    }
302
303    /// Returns the [TypeId] associated with this entity.
304    #[inline]
305    pub fn entity_type(&self) -> TypeId {
306        self.entity_type
307    }
308
309    /// Converts this entity handle into a weak variant, which does not prevent it from being released.
310    pub fn downgrade(&self) -> AnyWeakEntity {
311        AnyWeakEntity {
312            entity_id: self.entity_id,
313            entity_type: self.entity_type,
314            entity_ref_counts: self.entity_map.clone(),
315        }
316    }
317
318    /// Converts this entity handle into a strongly-typed entity handle of the given type.
319    /// If this entity handle is not of the specified type, returns itself as an error variant.
320    pub fn downcast<T: 'static>(self) -> Result<Entity<T>, AnyEntity> {
321        if TypeId::of::<T>() == self.entity_type {
322            Ok(Entity {
323                any_entity: self,
324                entity_type: PhantomData,
325            })
326        } else {
327            Err(self)
328        }
329    }
330}
331
332impl Clone for AnyEntity {
333    fn clone(&self) -> Self {
334        if let Some(entity_map) = self.entity_map.upgrade() {
335            let entity_map = entity_map.read();
336            let count = entity_map
337                .counts
338                .get(self.entity_id)
339                .expect("detected over-release of a entity");
340            let prev_count = count.fetch_add(1, SeqCst);
341            assert_ne!(prev_count, 0, "Detected over-release of a entity.");
342        }
343
344        Self {
345            entity_id: self.entity_id,
346            entity_type: self.entity_type,
347            entity_map: self.entity_map.clone(),
348            #[cfg(any(test, gpui_leak_detection))]
349            handle_id: self
350                .entity_map
351                .upgrade()
352                .unwrap()
353                .write()
354                .leak_detector
355                .handle_created(self.entity_id, None),
356        }
357    }
358}
359
360impl Drop for AnyEntity {
361    fn drop(&mut self) {
362        if let Some(entity_map) = self.entity_map.upgrade() {
363            let entity_map = entity_map.upgradable_read();
364            let count = entity_map
365                .counts
366                .get(self.entity_id)
367                .expect("detected over-release of a handle.");
368            let prev_count = count.fetch_sub(1, SeqCst);
369            assert_ne!(prev_count, 0, "Detected over-release of a entity.");
370            if prev_count == 1 {
371                // We were the last reference to this entity, so we can remove it.
372                let mut entity_map = RwLockUpgradableReadGuard::upgrade(entity_map);
373                entity_map.dropped_entity_ids.push(self.entity_id);
374            }
375        }
376
377        #[cfg(any(test, gpui_leak_detection))]
378        if let Some(entity_map) = self.entity_map.upgrade() {
379            entity_map
380                .write()
381                .leak_detector
382                .handle_released(self.entity_id, self.handle_id)
383        }
384    }
385}
386
387impl<T> From<Entity<T>> for AnyEntity {
388    #[inline]
389    fn from(entity: Entity<T>) -> Self {
390        entity.any_entity
391    }
392}
393
394impl Hash for AnyEntity {
395    #[inline]
396    fn hash<H: Hasher>(&self, state: &mut H) {
397        self.entity_id.hash(state);
398    }
399}
400
401impl PartialEq for AnyEntity {
402    #[inline]
403    fn eq(&self, other: &Self) -> bool {
404        self.entity_id == other.entity_id
405    }
406}
407
408impl Eq for AnyEntity {}
409
410impl Ord for AnyEntity {
411    #[inline]
412    fn cmp(&self, other: &Self) -> Ordering {
413        self.entity_id.cmp(&other.entity_id)
414    }
415}
416
417impl PartialOrd for AnyEntity {
418    #[inline]
419    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
420        Some(self.cmp(other))
421    }
422}
423
424impl std::fmt::Debug for AnyEntity {
425    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
426        f.debug_struct("AnyEntity")
427            .field("entity_id", &self.entity_id.as_u64())
428            .finish()
429    }
430}
431
432/// A strong, well-typed reference to a struct which is managed
433/// by GPUI
434#[derive(Deref, DerefMut)]
435pub struct Entity<T> {
436    #[deref]
437    #[deref_mut]
438    pub(crate) any_entity: AnyEntity,
439    pub(crate) entity_type: PhantomData<fn(T) -> T>,
440}
441
442impl<T> Sealed for Entity<T> {}
443
444impl<T: 'static> Entity<T> {
445    #[inline]
446    fn new(id: EntityId, entity_map: Weak<RwLock<EntityRefCounts>>) -> Self
447    where
448        T: 'static,
449    {
450        Self {
451            any_entity: AnyEntity::new(
452                id,
453                TypeId::of::<T>(),
454                entity_map,
455                #[cfg(any(test, gpui_leak_detection))]
456                std::any::type_name::<T>(),
457            ),
458            entity_type: PhantomData,
459        }
460    }
461
462    /// Get the entity ID associated with this entity
463    #[inline]
464    pub fn entity_id(&self) -> EntityId {
465        self.any_entity.entity_id
466    }
467
468    /// Downgrade this entity pointer to a non-retaining weak pointer
469    #[inline]
470    pub fn downgrade(&self) -> WeakEntity<T> {
471        WeakEntity {
472            any_entity: self.any_entity.downgrade(),
473            entity_type: self.entity_type,
474        }
475    }
476
477    /// Convert this into a dynamically typed entity.
478    #[inline]
479    pub fn into_any(self) -> AnyEntity {
480        self.any_entity
481    }
482
483    /// Grab a reference to this entity from the context.
484    #[inline]
485    pub fn read<'a>(&self, cx: &'a App) -> &'a T {
486        cx.entities.read(self)
487    }
488
489    /// Read the entity referenced by this handle with the given function.
490    #[inline]
491    pub fn read_with<R, C: AppContext>(&self, cx: &C, f: impl FnOnce(&T, &App) -> R) -> R {
492        cx.read_entity(self, f)
493    }
494
495    /// Updates the entity referenced by this handle with the given function.
496    #[inline]
497    pub fn update<R, C: AppContext>(
498        &self,
499        cx: &mut C,
500        update: impl FnOnce(&mut T, &mut Context<T>) -> R,
501    ) -> R {
502        cx.update_entity(self, update)
503    }
504
505    /// Updates the entity referenced by this handle with the given function.
506    #[inline]
507    pub fn as_mut<'a, C: AppContext>(&self, cx: &'a mut C) -> GpuiBorrow<'a, T> {
508        cx.as_mut(self)
509    }
510
511    /// Updates the entity referenced by this handle with the given function.
512    pub fn write<C: AppContext>(&self, cx: &mut C, value: T) {
513        self.update(cx, |entity, cx| {
514            *entity = value;
515            cx.notify();
516        })
517    }
518
519    /// Updates the entity referenced by this handle with the given function if
520    /// the referenced entity still exists, within a visual context that has a window.
521    /// Returns an error if the window has been closed.
522    #[inline]
523    pub fn update_in<R, C: VisualContext>(
524        &self,
525        cx: &mut C,
526        update: impl FnOnce(&mut T, &mut Window, &mut Context<T>) -> R,
527    ) -> C::Result<R> {
528        cx.update_window_entity(self, update)
529    }
530}
531
532impl<T> Clone for Entity<T> {
533    #[inline]
534    fn clone(&self) -> Self {
535        Self {
536            any_entity: self.any_entity.clone(),
537            entity_type: self.entity_type,
538        }
539    }
540}
541
542impl<T> std::fmt::Debug for Entity<T> {
543    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
544        f.debug_struct("Entity")
545            .field("entity_id", &self.any_entity.entity_id)
546            .field("entity_type", &type_name::<T>())
547            .finish()
548    }
549}
550
551impl<T> Hash for Entity<T> {
552    #[inline]
553    fn hash<H: Hasher>(&self, state: &mut H) {
554        self.any_entity.hash(state);
555    }
556}
557
558impl<T> PartialEq for Entity<T> {
559    #[inline]
560    fn eq(&self, other: &Self) -> bool {
561        self.any_entity == other.any_entity
562    }
563}
564
565impl<T> Eq for Entity<T> {}
566
567impl<T> PartialEq<WeakEntity<T>> for Entity<T> {
568    #[inline]
569    fn eq(&self, other: &WeakEntity<T>) -> bool {
570        self.any_entity.entity_id() == other.entity_id()
571    }
572}
573
574impl<T: 'static> Ord for Entity<T> {
575    #[inline]
576    fn cmp(&self, other: &Self) -> std::cmp::Ordering {
577        self.entity_id().cmp(&other.entity_id())
578    }
579}
580
581impl<T: 'static> PartialOrd for Entity<T> {
582    #[inline]
583    fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
584        Some(self.cmp(other))
585    }
586}
587
588/// A type erased, weak reference to a entity.
589#[derive(Clone)]
590pub struct AnyWeakEntity {
591    pub(crate) entity_id: EntityId,
592    entity_type: TypeId,
593    entity_ref_counts: Weak<RwLock<EntityRefCounts>>,
594}
595
596impl AnyWeakEntity {
597    /// Get the entity ID associated with this weak reference.
598    #[inline]
599    pub fn entity_id(&self) -> EntityId {
600        self.entity_id
601    }
602
603    /// Check if this weak handle can be upgraded, or if the entity has already been dropped
604    pub fn is_upgradable(&self) -> bool {
605        let ref_count = self
606            .entity_ref_counts
607            .upgrade()
608            .and_then(|ref_counts| Some(ref_counts.read().counts.get(self.entity_id)?.load(SeqCst)))
609            .unwrap_or(0);
610        ref_count > 0
611    }
612
613    /// Upgrade this weak entity reference to a strong reference.
614    pub fn upgrade(&self) -> Option<AnyEntity> {
615        let ref_counts = &self.entity_ref_counts.upgrade()?;
616        let ref_counts = ref_counts.read();
617        let ref_count = ref_counts.counts.get(self.entity_id)?;
618
619        if atomic_incr_if_not_zero(ref_count) == 0 {
620            // entity_id is in dropped_entity_ids
621            return None;
622        }
623        drop(ref_counts);
624
625        Some(AnyEntity {
626            entity_id: self.entity_id,
627            entity_type: self.entity_type,
628            entity_map: self.entity_ref_counts.clone(),
629            #[cfg(any(test, gpui_leak_detection))]
630            handle_id: self
631                .entity_ref_counts
632                .upgrade()
633                .unwrap()
634                .write()
635                .leak_detector
636                .handle_created(self.entity_id, None),
637        })
638    }
639
640    /// Asserts that the entity referenced by this weak handle has been fully released.
641    ///
642    /// # Example
643    ///
644    /// ```ignore
645    /// let entity = cx.new(|_| MyEntity::new());
646    /// let weak = entity.downgrade();
647    /// drop(entity);
648    ///
649    /// // Verify the entity was released
650    /// weak.assert_released();
651    /// ```
652    ///
653    /// # Debugging Leaks
654    ///
655    /// With leak detection enabled (build with `GPUI_LEAK_DETECTION=1`), a failure
656    /// lists each leaked handle. Also set the `LEAK_BACKTRACE` environment variable
657    /// to see where they were allocated:
658    ///
659    /// ```bash
660    /// GPUI_LEAK_DETECTION=1 LEAK_BACKTRACE=1 cargo test my_test
661    /// ```
662    ///
663    /// # Panics
664    ///
665    /// - Panics if any strong handles to the entity are still alive.
666    /// - Panics if the entity was recently dropped but cleanup hasn't completed yet
667    ///   (resources are retained until the end of the effect cycle).
668    #[cfg(any(test, feature = "test-support", gpui_leak_detection))]
669    pub fn assert_released(&self) {
670        #[cfg(any(test, gpui_leak_detection))]
671        self.entity_ref_counts
672            .upgrade()
673            .unwrap()
674            .write()
675            .leak_detector
676            .assert_released(self.entity_id);
677
678        if self
679            .entity_ref_counts
680            .upgrade()
681            .and_then(|ref_counts| Some(ref_counts.read().counts.get(self.entity_id)?.load(SeqCst)))
682            .is_some()
683        {
684            panic!(
685                "entity is still alive, or was recently dropped and its resources are retained \
686                 until the end of the effect cycle. Build with GPUI_LEAK_DETECTION=1 to list \
687                 its live handles."
688            )
689        }
690    }
691
692    /// Creates a weak entity that can never be upgraded.
693    pub fn new_invalid() -> Self {
694        /// To hold the invariant that all ids are unique, and considering that slotmap
695        /// increases their IDs from `0`, we can decrease ours from `u64::MAX` so these
696        /// two will never conflict (u64 is way too large).
697        static UNIQUE_NON_CONFLICTING_ID_GENERATOR: AtomicU64 = AtomicU64::new(u64::MAX);
698        let entity_id = UNIQUE_NON_CONFLICTING_ID_GENERATOR.fetch_sub(1, SeqCst);
699
700        Self {
701            // Safety:
702            //   Docs say this is safe but can be unspecified if slotmap changes the representation
703            //   after `1.0.7`, that said, providing a valid entity_id here is not necessary as long
704            //   as we guarantee that `entity_id` is never used if `entity_ref_counts` equals
705            //   to `Weak::new()` (that is, it's unable to upgrade), that is the invariant that
706            //   actually needs to be hold true.
707            //
708            //   And there is no sane reason to read an entity slot if `entity_ref_counts` can't be
709            //   read in the first place, so we're good!
710            entity_id: entity_id.into(),
711            entity_type: TypeId::of::<()>(),
712            entity_ref_counts: Weak::new(),
713        }
714    }
715}
716
717impl std::fmt::Debug for AnyWeakEntity {
718    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
719        f.debug_struct(type_name::<Self>())
720            .field("entity_id", &self.entity_id)
721            .field("entity_type", &self.entity_type)
722            .finish()
723    }
724}
725
726impl<T> From<WeakEntity<T>> for AnyWeakEntity {
727    #[inline]
728    fn from(entity: WeakEntity<T>) -> Self {
729        entity.any_entity
730    }
731}
732
733impl Hash for AnyWeakEntity {
734    #[inline]
735    fn hash<H: Hasher>(&self, state: &mut H) {
736        self.entity_id.hash(state);
737    }
738}
739
740impl PartialEq for AnyWeakEntity {
741    #[inline]
742    fn eq(&self, other: &Self) -> bool {
743        self.entity_id == other.entity_id
744    }
745}
746
747impl Eq for AnyWeakEntity {}
748
749impl Ord for AnyWeakEntity {
750    #[inline]
751    fn cmp(&self, other: &Self) -> Ordering {
752        self.entity_id.cmp(&other.entity_id)
753    }
754}
755
756impl PartialOrd for AnyWeakEntity {
757    #[inline]
758    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
759        Some(self.cmp(other))
760    }
761}
762
763/// A weak reference to a entity of the given type.
764#[derive(Deref, DerefMut)]
765pub struct WeakEntity<T> {
766    #[deref]
767    #[deref_mut]
768    any_entity: AnyWeakEntity,
769    entity_type: PhantomData<fn(T) -> T>,
770}
771
772impl<T> std::fmt::Debug for WeakEntity<T> {
773    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
774        f.debug_struct(type_name::<Self>())
775            .field("entity_id", &self.any_entity.entity_id)
776            .field("entity_type", &type_name::<T>())
777            .finish()
778    }
779}
780
781impl<T> Clone for WeakEntity<T> {
782    fn clone(&self) -> Self {
783        Self {
784            any_entity: self.any_entity.clone(),
785            entity_type: self.entity_type,
786        }
787    }
788}
789
790impl<T: 'static> WeakEntity<T> {
791    /// Upgrade this weak entity reference into a strong entity reference
792    pub fn upgrade(&self) -> Option<Entity<T>> {
793        Some(Entity {
794            any_entity: self.any_entity.upgrade()?,
795            entity_type: self.entity_type,
796        })
797    }
798
799    /// Updates the entity referenced by this handle with the given function if
800    /// the referenced entity still exists. Returns an error if the entity has
801    /// been released.
802    #[inline(always)]
803    pub fn update<C, R>(
804        &self,
805        cx: &mut C,
806        update: impl FnOnce(&mut T, &mut Context<T>) -> R,
807    ) -> Result<R>
808    where
809        C: AppContext,
810    {
811        let entity = self.upgrade().context("entity released")?;
812        Ok(cx.update_entity(&entity, update))
813    }
814
815    /// Updates the entity referenced by this handle with the given function if
816    /// the referenced entity still exists, within a visual context that has a window.
817    /// Returns an error if the entity has been released.
818    #[inline(always)]
819    pub fn update_in<C, R>(
820        &self,
821        cx: &mut C,
822        update: impl FnOnce(&mut T, &mut Window, &mut Context<T>) -> R,
823    ) -> Result<R>
824    where
825        C: AppContext,
826    {
827        let entity = self.upgrade().context("entity released")?;
828        cx.with_window(entity.entity_id(), |window, app| {
829            entity.update(app, |entity, cx| update(entity, window, cx))
830        })
831        .context("entity has no current window")
832    }
833
834    /// Reads the entity referenced by this handle with the given function if
835    /// the referenced entity still exists. Returns an error if the entity has
836    /// been released.
837    #[inline(always)]
838    pub fn read_with<C, R>(&self, cx: &C, read: impl FnOnce(&T, &App) -> R) -> Result<R>
839    where
840        C: AppContext,
841    {
842        let entity = self.upgrade().context("entity released")?;
843        Ok(cx.read_entity(&entity, read))
844    }
845
846    /// Create a new weak entity that can never be upgraded.
847    #[inline]
848    pub fn new_invalid() -> Self {
849        Self {
850            any_entity: AnyWeakEntity::new_invalid(),
851            entity_type: PhantomData,
852        }
853    }
854}
855
856impl<T> Hash for WeakEntity<T> {
857    #[inline]
858    fn hash<H: Hasher>(&self, state: &mut H) {
859        self.any_entity.hash(state);
860    }
861}
862
863impl<T> PartialEq for WeakEntity<T> {
864    #[inline]
865    fn eq(&self, other: &Self) -> bool {
866        self.any_entity == other.any_entity
867    }
868}
869
870impl<T> Eq for WeakEntity<T> {}
871
872impl<T> PartialEq<Entity<T>> for WeakEntity<T> {
873    #[inline]
874    fn eq(&self, other: &Entity<T>) -> bool {
875        self.entity_id() == other.any_entity.entity_id()
876    }
877}
878
879impl<T: 'static> Ord for WeakEntity<T> {
880    #[inline]
881    fn cmp(&self, other: &Self) -> Ordering {
882        self.entity_id().cmp(&other.entity_id())
883    }
884}
885
886impl<T: 'static> PartialOrd for WeakEntity<T> {
887    #[inline]
888    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
889        Some(self.cmp(other))
890    }
891}
892
893/// Controls whether backtraces are captured when entity handles are created.
894///
895/// Set the `LEAK_BACKTRACE` environment variable to any non-empty value to enable
896/// backtrace capture. This helps identify where leaked handles were allocated.
897#[cfg(any(test, gpui_leak_detection))]
898static LEAK_BACKTRACE: std::sync::LazyLock<bool> =
899    std::sync::LazyLock::new(|| std::env::var("LEAK_BACKTRACE").is_ok_and(|b| !b.is_empty()));
900
901/// Unique identifier for a specific entity handle instance.
902///
903/// This is distinct from `EntityId` - while multiple handles can point to the same
904/// entity (same `EntityId`), each handle has its own unique `HandleId`.
905#[cfg(any(test, gpui_leak_detection))]
906#[derive(Clone, Copy, Debug, Default, Hash, PartialEq, Eq)]
907pub(crate) struct HandleId {
908    id: u64,
909}
910
911/// Tracks entity handle allocations to detect leaks.
912///
913/// The leak detector is enabled in tests and when the `leak-detection` feature is active.
914/// It tracks every `Entity<T>` and `AnyEntity` handle that is created and released,
915/// allowing you to verify that all handles to an entity have been properly dropped.
916///
917/// # How do leaks happen?
918///
919/// Entities are reference-counted structures that can own other entities
920/// allowing to form cycles. If such a strong-reference counted cycle is
921/// created, all participating strong entities in this cycle will effectively
922/// leak as they cannot be released anymore.
923///
924/// Cycles can also happen if an entity owns a task or subscription that it
925/// itself owns a strong reference to the entity again.
926///
927/// # Usage
928///
929/// You can use `WeakEntity::assert_released` or `AnyWeakEntity::assert_released`
930/// to verify that an entity has been fully released:
931///
932/// ```ignore
933/// let entity = cx.new(|_| MyEntity::new());
934/// let weak = entity.downgrade();
935/// drop(entity);
936///
937/// // This will panic if any handles to the entity are still alive
938/// weak.assert_released();
939/// ```
940///
941/// # Debugging Leaks
942///
943/// When a leak is detected, the detector will panic with information about the leaked
944/// handles. To see where the leaked handles were allocated, set the `LEAK_BACKTRACE`
945/// environment variable:
946///
947/// ```bash
948/// LEAK_BACKTRACE=1 cargo test my_test
949/// ```
950///
951/// This will capture and display backtraces for each leaked handle, helping you
952/// identify where leaked handles were created.
953///
954/// # How It Works
955///
956/// - When an entity handle is created (via `Entity::new`, `Entity::clone`, or
957///   `WeakEntity::upgrade`), `handle_created` is called to register the handle.
958/// - When a handle is dropped, `handle_released` removes it from tracking.
959/// - `assert_released` verifies that no handles remain for a given entity.
960#[cfg(any(test, gpui_leak_detection))]
961pub(crate) struct LeakDetector {
962    next_handle_id: u64,
963    entity_handles: HashMap<EntityId, EntityLeakData>,
964}
965
966/// A snapshot of the set of alive entities at a point in time.
967///
968/// Created by [`LeakDetector::snapshot`]. Can later be passed to
969/// [`LeakDetector::assert_no_new_leaks`] to verify that no new entity
970/// handles remain between the snapshot and the current state.
971#[cfg(any(test, feature = "test-support", gpui_leak_detection))]
972#[derive(Default)]
973pub struct LeakDetectorSnapshot {
974    #[cfg(any(test, gpui_leak_detection))]
975    entity_ids: collections::HashSet<EntityId>,
976}
977
978#[cfg(any(test, gpui_leak_detection))]
979struct EntityLeakData {
980    handles: HashMap<HandleId, Option<backtrace::Backtrace>>,
981    type_name: &'static str,
982}
983
984#[cfg(any(test, gpui_leak_detection))]
985impl LeakDetector {
986    /// Records that a new handle has been created for the given entity.
987    ///
988    /// Returns a unique `HandleId` that must be passed to `handle_released` when
989    /// the handle is dropped. If `LEAK_BACKTRACE` is set, captures a backtrace
990    /// at the allocation site.
991    #[track_caller]
992    pub fn handle_created(
993        &mut self,
994        entity_id: EntityId,
995        type_name: Option<&'static str>,
996    ) -> HandleId {
997        let id = gpui_util::post_inc(&mut self.next_handle_id);
998        let handle_id = HandleId { id };
999        let handles = self
1000            .entity_handles
1001            .entry(entity_id)
1002            .or_insert_with(|| EntityLeakData {
1003                handles: HashMap::default(),
1004                type_name: type_name.unwrap_or("<unknown>"),
1005            });
1006        handles.handles.insert(
1007            handle_id,
1008            LEAK_BACKTRACE.then(backtrace::Backtrace::new_unresolved),
1009        );
1010        handle_id
1011    }
1012
1013    /// Records that a handle has been released (dropped).
1014    ///
1015    /// This removes the handle from tracking. The `handle_id` should be the same
1016    /// one returned by `handle_created` when the handle was allocated.
1017    pub fn handle_released(&mut self, entity_id: EntityId, handle_id: HandleId) {
1018        if let std::collections::hash_map::Entry::Occupied(mut data) =
1019            self.entity_handles.entry(entity_id)
1020        {
1021            data.get_mut().handles.remove(&handle_id);
1022            if data.get().handles.is_empty() {
1023                data.remove();
1024            }
1025        }
1026    }
1027
1028    /// Asserts that all handles to the given entity have been released.
1029    ///
1030    /// # Panics
1031    ///
1032    /// Panics if any handles to the entity are still alive. The panic message
1033    /// includes backtraces for each leaked handle if `LEAK_BACKTRACE` is set,
1034    /// otherwise it suggests setting the environment variable to get more info.
1035    pub fn assert_released(&mut self, entity_id: EntityId) {
1036        use std::fmt::Write as _;
1037
1038        if let Some(data) = self.entity_handles.remove(&entity_id) {
1039            let mut out = String::new();
1040            for (_, backtrace) in data.handles {
1041                if let Some(mut backtrace) = backtrace {
1042                    backtrace.resolve();
1043                    let backtrace = BacktraceFormatter(backtrace);
1044                    writeln!(out, "Leaked handle:\n{:?}", backtrace).unwrap();
1045                } else {
1046                    writeln!(
1047                        out,
1048                        "Leaked handle: (export LEAK_BACKTRACE to find allocation site)"
1049                    )
1050                    .unwrap();
1051                }
1052            }
1053            panic!("Handles for {} leaked:\n{out}", data.type_name);
1054        }
1055    }
1056
1057    /// Captures a snapshot of all entity IDs that currently have alive handles.
1058    ///
1059    /// The returned [`LeakDetectorSnapshot`] can later be passed to
1060    /// [`assert_no_new_leaks`](Self::assert_no_new_leaks) to verify that no
1061    /// entities created after the snapshot are still alive.
1062    pub fn snapshot(&self) -> LeakDetectorSnapshot {
1063        LeakDetectorSnapshot {
1064            entity_ids: self.entity_handles.keys().copied().collect(),
1065        }
1066    }
1067
1068    /// Asserts that no entities created after `snapshot` still have alive handles.
1069    ///
1070    /// Entities that were already tracked at the time of the snapshot are ignored,
1071    /// even if they still have handles. Only *new* entities (those whose
1072    /// `EntityId` was not present in the snapshot) are considered leaks.
1073    ///
1074    /// # Panics
1075    ///
1076    /// Panics if any new entity handles exist. The panic message lists every
1077    /// leaked entity with its type name, and includes allocation-site backtraces
1078    /// when `LEAK_BACKTRACE` is set.
1079    pub fn assert_no_new_leaks(&self, snapshot: &LeakDetectorSnapshot) {
1080        use std::fmt::Write as _;
1081
1082        let mut out = String::new();
1083        for (entity_id, data) in &self.entity_handles {
1084            if snapshot.entity_ids.contains(entity_id) {
1085                continue;
1086            }
1087            for (_, backtrace) in &data.handles {
1088                if let Some(backtrace) = backtrace {
1089                    let mut backtrace = backtrace.clone();
1090                    backtrace.resolve();
1091                    let backtrace = BacktraceFormatter(backtrace);
1092                    writeln!(
1093                        out,
1094                        "Leaked handle for entity {} ({entity_id:?}):\n{:?}",
1095                        data.type_name, backtrace
1096                    )
1097                    .unwrap();
1098                } else {
1099                    writeln!(
1100                        out,
1101                        "Leaked handle for entity {} ({entity_id:?}): (export LEAK_BACKTRACE to find allocation site)",
1102                        data.type_name
1103                    )
1104                    .unwrap();
1105                }
1106            }
1107        }
1108
1109        if !out.is_empty() {
1110            panic!("New entity leaks detected since snapshot:\n{out}");
1111        }
1112    }
1113}
1114
1115#[cfg(any(test, gpui_leak_detection))]
1116impl Drop for LeakDetector {
1117    fn drop(&mut self) {
1118        use std::fmt::Write;
1119
1120        if self.entity_handles.is_empty() || std::thread::panicking() {
1121            return;
1122        }
1123
1124        let mut out = String::new();
1125        for (entity_id, data) in self.entity_handles.drain() {
1126            for (_handle, backtrace) in data.handles {
1127                if let Some(mut backtrace) = backtrace {
1128                    backtrace.resolve();
1129                    let backtrace = BacktraceFormatter(backtrace);
1130                    writeln!(
1131                        out,
1132                        "Leaked handle for entity {} ({entity_id:?}):\n{:?}",
1133                        data.type_name, backtrace
1134                    )
1135                    .unwrap();
1136                } else {
1137                    writeln!(
1138                        out,
1139                        "Leaked handle for entity {} ({entity_id:?}): (export LEAK_BACKTRACE to find allocation site)",
1140                        data.type_name
1141                    )
1142                    .unwrap();
1143                }
1144            }
1145        }
1146        panic!("Exited with leaked handles:\n{out}");
1147    }
1148}
1149
1150#[cfg(any(test, gpui_leak_detection))]
1151struct BacktraceFormatter(backtrace::Backtrace);
1152
1153#[cfg(any(test, gpui_leak_detection))]
1154impl fmt::Debug for BacktraceFormatter {
1155    fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
1156        use backtrace::{BacktraceFmt, BytesOrWideString, PrintFmt};
1157
1158        let style = if fmt.alternate() {
1159            PrintFmt::Full
1160        } else {
1161            PrintFmt::Short
1162        };
1163
1164        // When printing paths we try to strip the cwd if it exists, otherwise
1165        // we just print the path as-is. Note that we also only do this for the
1166        // short format, because if it's full we presumably want to print
1167        // everything.
1168        let cwd = std::env::current_dir();
1169        let mut print_path = move |fmt: &mut fmt::Formatter<'_>, path: BytesOrWideString<'_>| {
1170            let path = path.into_path_buf();
1171            if style != PrintFmt::Full {
1172                if let Ok(cwd) = &cwd {
1173                    if let Ok(suffix) = path.strip_prefix(cwd) {
1174                        return fmt::Display::fmt(&suffix.display(), fmt);
1175                    }
1176                }
1177            }
1178            fmt::Display::fmt(&path.display(), fmt)
1179        };
1180
1181        let mut f = BacktraceFmt::new(fmt, style, &mut print_path);
1182        f.add_context()?;
1183        let mut strip = true;
1184        for frame in self.0.frames() {
1185            if let [symbol, ..] = frame.symbols()
1186                && let Some(name) = symbol.name()
1187                && let Some(filename) = name.as_str()
1188            {
1189                match filename {
1190                    "test::run_test_in_process"
1191                    | "scheduler::executor::spawn_local_with_source_location::impl$1::poll<core::pin::Pin<alloc::boxed::Box<dyn$<core::future::future::Future<assoc$<Output,enum2$<core::result::Result<workspace::OpenResult,anyhow::Error> > > > >,alloc::alloc::Global> > >" => {
1192                        strip = true
1193                    }
1194                    "gpui::app::entity_map::LeakDetector::handle_created" => {
1195                        strip = false;
1196                        continue;
1197                    }
1198                    "zed::main" => {
1199                        strip = true;
1200                        f.frame().backtrace_frame(frame)?;
1201                    }
1202                    _ => {}
1203                }
1204            }
1205            if strip {
1206                continue;
1207            }
1208            f.frame().backtrace_frame(frame)?;
1209        }
1210        f.finish()?;
1211        Ok(())
1212    }
1213}
1214
1215#[cfg(test)]
1216mod test {
1217    use crate::EntityMap;
1218
1219    struct TestEntity {
1220        pub i: i32,
1221    }
1222
1223    #[test]
1224    fn test_entity_map_slot_assignment_before_cleanup() {
1225        // Tests that slots are not re-used before take_dropped.
1226        let mut entity_map = EntityMap::new();
1227
1228        let slot = entity_map.reserve::<TestEntity>();
1229        entity_map.insert(slot, TestEntity { i: 1 });
1230
1231        let slot = entity_map.reserve::<TestEntity>();
1232        entity_map.insert(slot, TestEntity { i: 2 });
1233
1234        let dropped = entity_map.take_dropped();
1235        assert_eq!(dropped.len(), 2);
1236
1237        assert_eq!(
1238            dropped
1239                .into_iter()
1240                .map(|(_, entity)| entity.downcast::<TestEntity>().unwrap().i)
1241                .collect::<Vec<i32>>(),
1242            vec![1, 2],
1243        );
1244    }
1245
1246    #[test]
1247    fn test_entity_map_weak_upgrade_before_cleanup() {
1248        // Tests that weak handles are not upgraded before take_dropped
1249        let mut entity_map = EntityMap::new();
1250
1251        let slot = entity_map.reserve::<TestEntity>();
1252        let handle = entity_map.insert(slot, TestEntity { i: 1 });
1253        let weak = handle.downgrade();
1254        drop(handle);
1255
1256        let strong = weak.upgrade();
1257        assert_eq!(strong, None);
1258
1259        let dropped = entity_map.take_dropped();
1260        assert_eq!(dropped.len(), 1);
1261
1262        assert_eq!(
1263            dropped
1264                .into_iter()
1265                .map(|(_, entity)| entity.downcast::<TestEntity>().unwrap().i)
1266                .collect::<Vec<i32>>(),
1267            vec![1],
1268        );
1269    }
1270
1271    #[test]
1272    fn test_leak_detector_snapshot_no_leaks() {
1273        let mut entity_map = EntityMap::new();
1274
1275        let slot = entity_map.reserve::<TestEntity>();
1276        let pre_existing = entity_map.insert(slot, TestEntity { i: 1 });
1277
1278        let snapshot = entity_map.leak_detector_snapshot();
1279
1280        let slot = entity_map.reserve::<TestEntity>();
1281        let temporary = entity_map.insert(slot, TestEntity { i: 2 });
1282        drop(temporary);
1283
1284        entity_map.assert_no_new_leaks(&snapshot);
1285
1286        drop(pre_existing);
1287    }
1288
1289    #[test]
1290    #[should_panic(expected = "New entity leaks detected since snapshot")]
1291    fn test_leak_detector_snapshot_detects_new_leak() {
1292        let mut entity_map = EntityMap::new();
1293
1294        let slot = entity_map.reserve::<TestEntity>();
1295        let pre_existing = entity_map.insert(slot, TestEntity { i: 1 });
1296
1297        let snapshot = entity_map.leak_detector_snapshot();
1298
1299        let slot = entity_map.reserve::<TestEntity>();
1300        let leaked = entity_map.insert(slot, TestEntity { i: 2 });
1301
1302        // `leaked` is still alive, so this should panic.
1303        entity_map.assert_no_new_leaks(&snapshot);
1304
1305        drop(pre_existing);
1306        drop(leaked);
1307    }
1308}