Skip to main content

gpui_wgpu/
wgpu_atlas.rs

1use anyhow::{Context as _, Result};
2use etagere::{BucketedAtlasAllocator, size2};
3use gpui::{
4    AtlasBackend, AtlasKey, AtlasState, AtlasTextureId, AtlasTextureKind, AtlasTextureList,
5    AtlasTile, Bounds, DevicePixels, PlatformAtlas, Point, Size,
6};
7use parking_lot::Mutex;
8use std::{borrow::Cow, ops, sync::Arc};
9
10use crate::WgpuContext;
11
12fn device_size_to_etagere(size: Size<DevicePixels>) -> etagere::Size {
13    size2(size.width.0, size.height.0)
14}
15
16fn etagere_point_to_device(point: etagere::Point) -> Point<DevicePixels> {
17    Point {
18        x: DevicePixels(point.x),
19        y: DevicePixels(point.y),
20    }
21}
22
23pub struct WgpuAtlas(Mutex<AtlasState<WgpuAtlasTextures>>);
24
25struct PendingUpload {
26    id: AtlasTextureId,
27    bounds: Bounds<DevicePixels>,
28    data: Vec<u8>,
29}
30
31struct WgpuAtlasTextures {
32    device: Arc<wgpu::Device>,
33    queue: Arc<wgpu::Queue>,
34    max_texture_size: u32,
35    color_texture_format: wgpu::TextureFormat,
36    storage: WgpuAtlasStorage,
37    pending_uploads: Vec<PendingUpload>,
38    next_texture_generation: u64,
39}
40
41pub struct WgpuTextureInfo {
42    pub view: wgpu::TextureView,
43    /// Distinguishes new textures that reuse an [`AtlasTextureId`].
44    pub generation: u64,
45}
46
47impl WgpuAtlas {
48    pub fn new(
49        device: Arc<wgpu::Device>,
50        queue: Arc<wgpu::Queue>,
51        color_texture_format: wgpu::TextureFormat,
52    ) -> Self {
53        let max_texture_size = device.limits().max_texture_dimension_2d;
54        WgpuAtlas(Mutex::new(AtlasState::new(WgpuAtlasTextures {
55            device,
56            queue,
57            max_texture_size,
58            color_texture_format,
59            storage: WgpuAtlasStorage::default(),
60            pending_uploads: Vec::new(),
61            next_texture_generation: 0,
62        })))
63    }
64
65    pub fn from_context(context: &WgpuContext) -> Self {
66        Self::new(
67            context.device.clone(),
68            context.queue.clone(),
69            context.color_texture_format(),
70        )
71    }
72
73    pub fn before_frame(&self) {
74        let mut lock = self.0.lock();
75        lock.backend.flush_uploads();
76    }
77
78    /// Returns the view backing `id`, or `None` once every tile in it has been
79    /// removed. A scene can still reference such a texture when a cached view
80    /// replays a paint from before the image was dropped, so callers must skip
81    /// those sprites rather than assume the texture exists.
82    pub fn get_texture_info(&self, id: AtlasTextureId) -> Option<WgpuTextureInfo> {
83        let lock = self.0.lock();
84        let texture = lock.backend.storage.get(id)?;
85        Some(WgpuTextureInfo {
86            view: texture.view.clone(),
87            generation: texture.generation,
88        })
89    }
90
91    /// Clears all cached textures and tiles, forcing them to be recreated.
92    /// Use this for incremental recovery when the device is still valid.
93    pub fn clear(&self) {
94        self.0.lock().clear(|textures| {
95            textures.storage = WgpuAtlasStorage::default();
96            textures.pending_uploads.clear();
97        });
98    }
99
100    /// Handles device lost by clearing all textures and cached tiles.
101    /// The atlas will lazily recreate textures as needed on subsequent frames.
102    pub fn handle_device_lost(&self, context: &WgpuContext) {
103        self.0.lock().clear(|textures| {
104            textures.device = context.device.clone();
105            textures.queue = context.queue.clone();
106            textures.color_texture_format = context.color_texture_format();
107            textures.storage = WgpuAtlasStorage::default();
108            textures.pending_uploads.clear();
109        });
110    }
111}
112
113impl PlatformAtlas for WgpuAtlas {
114    fn get_or_insert_with<'a>(
115        &self,
116        key: AtlasKey,
117        build: &mut dyn FnMut() -> Result<Option<(Size<DevicePixels>, Cow<'a, [u8]>)>>,
118    ) -> Result<Option<AtlasTile>> {
119        self.0.lock().get_or_insert_with(key, build)
120    }
121
122    fn remove(&self, key: &AtlasKey) {
123        self.0.lock().remove(key);
124    }
125}
126
127impl AtlasBackend for WgpuAtlasTextures {
128    fn insert(
129        &mut self,
130        kind: AtlasTextureKind,
131        size: Size<DevicePixels>,
132        bytes: &[u8],
133    ) -> Result<AtlasTile> {
134        let tile = self.allocate(size, kind).context("failed to allocate")?;
135        self.upload_texture(tile.texture_id, tile.bounds, bytes);
136        Ok(tile)
137    }
138
139    fn remove(&mut self, tile: AtlasTile) {
140        let id = tile.texture_id;
141        let Some(texture_slot) = self.storage[id.kind].textures.get_mut(id.index as usize) else {
142            return;
143        };
144
145        if let Some(mut texture) = texture_slot.take() {
146            texture.allocator.deallocate(tile.tile_id.into());
147            texture.decrement_ref_count();
148            if texture.is_unreferenced() {
149                self.pending_uploads
150                    .retain(|upload| upload.id != texture.id);
151                self.storage[id.kind]
152                    .free_list
153                    .push(texture.id.index as usize);
154            } else {
155                *texture_slot = Some(texture);
156            }
157        }
158    }
159}
160
161impl WgpuAtlasTextures {
162    fn allocate(
163        &mut self,
164        size: Size<DevicePixels>,
165        texture_kind: AtlasTextureKind,
166    ) -> Option<AtlasTile> {
167        {
168            let textures = &mut self.storage[texture_kind];
169
170            if let Some(tile) = textures
171                .iter_mut()
172                .rev()
173                .find_map(|texture| texture.allocate(size))
174            {
175                return Some(tile);
176            }
177        }
178
179        let texture = self.push_texture(size, texture_kind);
180        texture.allocate(size)
181    }
182
183    fn push_texture(
184        &mut self,
185        min_size: Size<DevicePixels>,
186        kind: AtlasTextureKind,
187    ) -> &mut WgpuAtlasTexture {
188        const DEFAULT_ATLAS_SIZE: Size<DevicePixels> = Size {
189            width: DevicePixels(1024),
190            height: DevicePixels(1024),
191        };
192        let max_texture_size = self.max_texture_size as i32;
193        let max_atlas_size = Size {
194            width: DevicePixels(max_texture_size),
195            height: DevicePixels(max_texture_size),
196        };
197
198        let size = min_size.min(&max_atlas_size).max(&DEFAULT_ATLAS_SIZE);
199        let format = match kind {
200            AtlasTextureKind::Monochrome => wgpu::TextureFormat::R8Unorm,
201            AtlasTextureKind::Subpixel | AtlasTextureKind::Polychrome => self.color_texture_format,
202        };
203
204        let texture = self.device.create_texture(&wgpu::TextureDescriptor {
205            label: Some("atlas"),
206            size: wgpu::Extent3d {
207                width: size.width.0 as u32,
208                height: size.height.0 as u32,
209                depth_or_array_layers: 1,
210            },
211            mip_level_count: 1,
212            sample_count: 1,
213            dimension: wgpu::TextureDimension::D2,
214            format,
215            usage: wgpu::TextureUsages::TEXTURE_BINDING | wgpu::TextureUsages::COPY_DST,
216            view_formats: &[],
217        });
218
219        let view = texture.create_view(&wgpu::TextureViewDescriptor::default());
220
221        let texture_list = &mut self.storage[kind];
222        let index = texture_list.free_list.pop();
223        let generation = self.next_texture_generation;
224        self.next_texture_generation = self.next_texture_generation.wrapping_add(1);
225
226        let atlas_texture = WgpuAtlasTexture {
227            id: AtlasTextureId {
228                index: index.unwrap_or(texture_list.textures.len()) as u32,
229                kind,
230            },
231            generation,
232            allocator: BucketedAtlasAllocator::new(device_size_to_etagere(size)),
233            format,
234            texture,
235            view,
236            live_atlas_keys: 0,
237        };
238
239        if let Some(ix) = index {
240            texture_list.textures[ix] = Some(atlas_texture);
241            texture_list
242                .textures
243                .get_mut(ix)
244                .and_then(|t| t.as_mut())
245                .expect("texture must exist")
246        } else {
247            texture_list.textures.push(Some(atlas_texture));
248            texture_list
249                .textures
250                .last_mut()
251                .and_then(|t| t.as_mut())
252                .expect("texture must exist")
253        }
254    }
255
256    fn upload_texture(&mut self, id: AtlasTextureId, bounds: Bounds<DevicePixels>, bytes: &[u8]) {
257        let data = self
258            .storage
259            .get(id)
260            .map(|texture| swizzle_upload_data(bytes, texture.format))
261            .unwrap_or_else(|| bytes.to_vec());
262
263        self.pending_uploads
264            .push(PendingUpload { id, bounds, data });
265    }
266
267    fn flush_uploads(&mut self) {
268        for upload in self.pending_uploads.drain(..) {
269            let Some(texture) = self.storage.get(upload.id) else {
270                continue;
271            };
272            let bytes_per_pixel = texture.bytes_per_pixel();
273
274            self.queue.write_texture(
275                wgpu::TexelCopyTextureInfo {
276                    texture: &texture.texture,
277                    mip_level: 0,
278                    origin: wgpu::Origin3d {
279                        x: upload.bounds.origin.x.0 as u32,
280                        y: upload.bounds.origin.y.0 as u32,
281                        z: 0,
282                    },
283                    aspect: wgpu::TextureAspect::All,
284                },
285                &upload.data,
286                wgpu::TexelCopyBufferLayout {
287                    offset: 0,
288                    bytes_per_row: Some(upload.bounds.size.width.0 as u32 * bytes_per_pixel as u32),
289                    rows_per_image: None,
290                },
291                wgpu::Extent3d {
292                    width: upload.bounds.size.width.0 as u32,
293                    height: upload.bounds.size.height.0 as u32,
294                    depth_or_array_layers: 1,
295                },
296            );
297        }
298    }
299}
300
301#[derive(Default)]
302struct WgpuAtlasStorage {
303    monochrome_textures: AtlasTextureList<WgpuAtlasTexture>,
304    subpixel_textures: AtlasTextureList<WgpuAtlasTexture>,
305    polychrome_textures: AtlasTextureList<WgpuAtlasTexture>,
306}
307
308impl ops::Index<AtlasTextureKind> for WgpuAtlasStorage {
309    type Output = AtlasTextureList<WgpuAtlasTexture>;
310    fn index(&self, kind: AtlasTextureKind) -> &Self::Output {
311        match kind {
312            AtlasTextureKind::Monochrome => &self.monochrome_textures,
313            AtlasTextureKind::Subpixel => &self.subpixel_textures,
314            AtlasTextureKind::Polychrome => &self.polychrome_textures,
315        }
316    }
317}
318
319impl ops::IndexMut<AtlasTextureKind> for WgpuAtlasStorage {
320    fn index_mut(&mut self, kind: AtlasTextureKind) -> &mut Self::Output {
321        match kind {
322            AtlasTextureKind::Monochrome => &mut self.monochrome_textures,
323            AtlasTextureKind::Subpixel => &mut self.subpixel_textures,
324            AtlasTextureKind::Polychrome => &mut self.polychrome_textures,
325        }
326    }
327}
328
329impl WgpuAtlasStorage {
330    fn get(&self, id: AtlasTextureId) -> Option<&WgpuAtlasTexture> {
331        self[id.kind]
332            .textures
333            .get(id.index as usize)
334            .and_then(|t| t.as_ref())
335    }
336}
337
338struct WgpuAtlasTexture {
339    id: AtlasTextureId,
340    generation: u64,
341    allocator: BucketedAtlasAllocator,
342    texture: wgpu::Texture,
343    view: wgpu::TextureView,
344    format: wgpu::TextureFormat,
345    live_atlas_keys: u32,
346}
347
348impl WgpuAtlasTexture {
349    fn allocate(&mut self, size: Size<DevicePixels>) -> Option<AtlasTile> {
350        let allocation = self.allocator.allocate(device_size_to_etagere(size))?;
351        let tile = AtlasTile {
352            texture_id: self.id,
353            tile_id: allocation.id.into(),
354            padding: 0,
355            bounds: Bounds {
356                origin: etagere_point_to_device(allocation.rectangle.min),
357                size,
358            },
359        };
360        self.live_atlas_keys += 1;
361        Some(tile)
362    }
363
364    fn bytes_per_pixel(&self) -> u8 {
365        match self.format {
366            wgpu::TextureFormat::R8Unorm => 1,
367            wgpu::TextureFormat::Bgra8Unorm | wgpu::TextureFormat::Rgba8Unorm => 4,
368            _ => 4,
369        }
370    }
371
372    fn decrement_ref_count(&mut self) {
373        self.live_atlas_keys -= 1;
374    }
375
376    fn is_unreferenced(&self) -> bool {
377        self.live_atlas_keys == 0
378    }
379}
380
381fn swizzle_upload_data(bytes: &[u8], format: wgpu::TextureFormat) -> Vec<u8> {
382    match format {
383        wgpu::TextureFormat::Rgba8Unorm => {
384            let mut data = bytes.to_vec();
385            for pixel in data.chunks_exact_mut(4) {
386                pixel.swap(0, 2);
387            }
388            data
389        }
390        _ => bytes.to_vec(),
391    }
392}
393
394#[cfg(all(test, not(target_family = "wasm")))]
395mod tests {
396    use super::*;
397    use gpui::block_on;
398    use gpui::{ImageId, RenderImageParams};
399    use std::sync::Arc;
400
401    fn test_device_and_queue() -> anyhow::Result<(Arc<wgpu::Device>, Arc<wgpu::Queue>)> {
402        block_on(async {
403            let instance = wgpu::Instance::new(wgpu::InstanceDescriptor {
404                backends: wgpu::Backends::all(),
405                flags: wgpu::InstanceFlags::default(),
406                backend_options: wgpu::BackendOptions::default(),
407                memory_budget_thresholds: wgpu::MemoryBudgetThresholds::default(),
408                display: None,
409            });
410            let adapter = instance
411                .request_adapter(&wgpu::RequestAdapterOptions {
412                    power_preference: wgpu::PowerPreference::LowPower,
413                    compatible_surface: None,
414                    force_fallback_adapter: false,
415                })
416                .await
417                .map_err(|error| anyhow::anyhow!("failed to request adapter: {error}"))?;
418            let (device, queue) = adapter
419                .request_device(&wgpu::DeviceDescriptor {
420                    label: Some("wgpu_atlas_test_device"),
421                    required_features: wgpu::Features::empty(),
422                    required_limits: wgpu::Limits::downlevel_defaults()
423                        .using_resolution(adapter.limits())
424                        .using_alignment(adapter.limits()),
425                    memory_hints: wgpu::MemoryHints::MemoryUsage,
426                    trace: wgpu::Trace::Off,
427                    experimental_features: wgpu::ExperimentalFeatures::disabled(),
428                })
429                .await
430                .map_err(|error| anyhow::anyhow!("failed to request device: {error}"))?;
431            Ok((Arc::new(device), Arc::new(queue)))
432        })
433    }
434
435    #[test]
436    fn before_frame_skips_uploads_for_removed_texture() -> anyhow::Result<()> {
437        let (device, queue) = test_device_and_queue()?;
438
439        let atlas = WgpuAtlas::new(device, queue, wgpu::TextureFormat::Bgra8Unorm);
440        let key = AtlasKey::Image(RenderImageParams {
441            image_id: ImageId(1),
442            frame_index: 0,
443        });
444        let size = Size {
445            width: DevicePixels(1),
446            height: DevicePixels(1),
447        };
448        let mut build = || Ok(Some((size, Cow::Owned(vec![0, 0, 0, 255]))));
449
450        // Regression test: before the fix, this panicked in flush_uploads
451        atlas
452            .get_or_insert_with(key.clone(), &mut build)?
453            .expect("tile should be created");
454        atlas.remove(&key);
455        atlas.before_frame();
456        Ok(())
457    }
458
459    #[test]
460    fn remove_deallocates_tile_space_for_reuse() -> anyhow::Result<()> {
461        let (device, queue) = test_device_and_queue()?;
462        let atlas = WgpuAtlas::new(device, queue, wgpu::TextureFormat::Bgra8Unorm);
463
464        let small = Size {
465            width: DevicePixels(64),
466            height: DevicePixels(64),
467        };
468        let big = Size {
469            width: DevicePixels(700),
470            height: DevicePixels(700),
471        };
472
473        let make_key = |image_id: usize| {
474            AtlasKey::Image(RenderImageParams {
475                image_id: ImageId(image_id),
476                frame_index: 0,
477            })
478        };
479        let insert = |key: AtlasKey, size: Size<DevicePixels>| {
480            let byte_count = (size.width.0 as usize) * (size.height.0 as usize) * 4;
481            atlas
482                .get_or_insert_with(key, &mut || {
483                    Ok(Some((size, Cow::Owned(vec![0u8; byte_count]))))
484                })
485                .expect("allocation should succeed")
486                .expect("callback returns Some")
487        };
488
489        let keeper_key = make_key(1);
490        let big_key_a = make_key(2);
491        let big_key_b = make_key(3);
492
493        let keeper_tile = insert(keeper_key, small);
494        let tile_a = insert(big_key_a.clone(), big);
495        assert_eq!(keeper_tile.texture_id, tile_a.texture_id);
496
497        atlas.remove(&big_key_a);
498        let tile_b = insert(big_key_b, big);
499        assert_eq!(tile_b.texture_id, keeper_tile.texture_id);
500        Ok(())
501    }
502
503    #[test]
504    fn reused_texture_id_has_new_generation() -> anyhow::Result<()> {
505        let (device, queue) = test_device_and_queue()?;
506        let atlas = WgpuAtlas::new(device, queue, wgpu::TextureFormat::Bgra8Unorm);
507        let size = Size {
508            width: DevicePixels(700),
509            height: DevicePixels(700),
510        };
511        let make_key = |image_id| {
512            AtlasKey::Image(RenderImageParams {
513                image_id: ImageId(image_id),
514                frame_index: 0,
515            })
516        };
517        let insert = |key: AtlasKey| {
518            atlas
519                .get_or_insert_with(key, &mut || {
520                    Ok(Some((
521                        size,
522                        Cow::Owned(vec![0; size.width.0 as usize * size.height.0 as usize * 4]),
523                    )))
524                })
525                .expect("allocation should succeed")
526                .expect("callback returns Some")
527        };
528
529        let first_key = make_key(1);
530        let first_tile = insert(first_key.clone());
531        let first_generation = atlas
532            .get_texture_info(first_tile.texture_id)
533            .context("first texture should exist")?
534            .generation;
535        atlas.remove(&first_key);
536        assert!(atlas.get_texture_info(first_tile.texture_id).is_none());
537
538        let second_tile = insert(make_key(2));
539        let second_generation = atlas
540            .get_texture_info(second_tile.texture_id)
541            .context("second texture should exist")?
542            .generation;
543
544        assert_eq!(second_tile.texture_id, first_tile.texture_id);
545        assert_ne!(second_generation, first_generation);
546        Ok(())
547    }
548
549    #[test]
550    fn swizzle_upload_data_preserves_bgra_uploads() {
551        let input = vec![0x10, 0x20, 0x30, 0x40];
552        assert_eq!(
553            swizzle_upload_data(&input, wgpu::TextureFormat::Bgra8Unorm),
554            input
555        );
556    }
557
558    #[test]
559    fn swizzle_upload_data_converts_bgra_to_rgba() {
560        let input = vec![0x10, 0x20, 0x30, 0x40, 0xAA, 0xBB, 0xCC, 0xDD];
561        assert_eq!(
562            swizzle_upload_data(&input, wgpu::TextureFormat::Rgba8Unorm),
563            vec![0x30, 0x20, 0x10, 0x40, 0xCC, 0xBB, 0xAA, 0xDD]
564        );
565    }
566}