wgpu_core/command/
query.rs

1use alloc::{sync::Arc, vec, vec::Vec};
2use bit_vec::BitVec;
3use core::{iter, mem};
4
5use crate::{
6    command::{encoder::EncodingState, ArcCommand, EncoderStateError, InnerCommandEncoder},
7    device::{Device, DeviceError, MissingFeatures},
8    init_tracker::MemoryInitKind,
9    resource::{
10        Buffer, DestroyedResourceError, InvalidResourceError, Labeled as _,
11        MissingBufferUsageError, ParentDevice, QuerySet, RawResourceAccess, Trackable,
12    },
13    snatch::SnatchGuard,
14    track::{QuerySetTracker, TrackerIndex},
15    FastHashMap,
16};
17use thiserror::Error;
18use wgt::{
19    error::{ErrorType, WebGpuError},
20    BufferAddress,
21};
22
23pub(crate) struct DeferredQuerySetResolve {
24    pub(crate) query_set: Arc<QuerySet>,
25    pub(crate) query_set_writes: Option<BitVec>,
26    pub(crate) start_query: u32,
27    pub(crate) end_query: u32,
28    pub(crate) dst_buffer: Arc<Buffer>,
29    pub(crate) destination_offset: BufferAddress,
30    /// Bytes per query slot in the destination buffer
31    /// (accounts for pipeline-statistics element count * `QUERY_SIZE`).
32    pub(crate) stride: u64,
33    /// Index into [`InnerCommandEncoder::list`] at which a new command buffer
34    /// for the resolve operation must be inserted at submit time so that it
35    /// executes at exactly the position it was recorded.
36    pub(crate) insertion_point: usize,
37}
38
39pub(crate) type QuerySetWrites = FastHashMap<TrackerIndex, BitVec>;
40
41pub(super) fn record_pass_timestamp_writes(
42    tw: &crate::command::PassTimestampWrites,
43    query_set_writes: &mut QuerySetWrites,
44) {
45    for index in tw
46        .beginning_of_pass_write_index
47        .into_iter()
48        .chain(tw.end_of_pass_write_index)
49    {
50        record_query_write(query_set_writes, &tw.query_set, index);
51    }
52}
53
54pub(crate) fn record_query_write(
55    query_set_writes: &mut QuerySetWrites,
56    query_set: &Arc<QuerySet>,
57    slot_index: u32,
58) {
59    query_set_writes
60        .entry(query_set.tracker_index())
61        .or_insert_with(|| BitVec::from_elem(query_set.desc.count as usize, false))
62        .set(slot_index as usize, true);
63}
64
65#[derive(Debug)]
66pub(crate) struct QueryResetMap {
67    map: FastHashMap<TrackerIndex, (Vec<bool>, Arc<QuerySet>)>,
68}
69impl QueryResetMap {
70    pub fn new() -> Self {
71        Self {
72            map: FastHashMap::default(),
73        }
74    }
75
76    pub fn use_query_set(&mut self, query_set: &Arc<QuerySet>, query: u32) -> bool {
77        let vec_pair = self
78            .map
79            .entry(query_set.tracker_index())
80            .or_insert_with(|| {
81                (
82                    vec![false; query_set.desc.count as usize],
83                    query_set.clone(),
84                )
85            });
86
87        mem::replace(&mut vec_pair.0[query as usize], true)
88    }
89
90    pub fn reset_queries(
91        &mut self,
92        raw_encoder: &mut dyn hal::DynCommandEncoder,
93        snatch_guard: &SnatchGuard<'_>,
94    ) -> Result<(), DestroyedResourceError> {
95        for (_, (state, query_set)) in self.map.drain() {
96            debug_assert_eq!(state.len(), query_set.desc.count as usize);
97
98            // Need to find all "runs" of values which need resets. If the state vector is:
99            // [false, true, true, false, true], we want to reset [1..3, 4..5]. This minimizes
100            // the amount of resets needed.
101            let mut run_start: Option<u32> = None;
102            for (idx, value) in state.into_iter().chain(iter::once(false)).enumerate() {
103                match (run_start, value) {
104                    // We're inside of a run, do nothing
105                    (Some(..), true) => {}
106                    // We've hit the end of a run, dispatch a reset
107                    (Some(start), false) => {
108                        run_start = None;
109                        unsafe {
110                            raw_encoder
111                                .reset_queries(query_set.try_raw(snatch_guard)?, start..idx as u32)
112                        };
113                    }
114                    // We're starting a run
115                    (None, true) => {
116                        run_start = Some(idx as u32);
117                    }
118                    // We're in a run of falses, do nothing.
119                    (None, false) => {}
120                }
121            }
122        }
123        Ok(())
124    }
125}
126
127#[derive(Debug, Copy, Clone, PartialEq, Eq)]
128pub enum SimplifiedQueryType {
129    Occlusion,
130    Timestamp,
131    PipelineStatistics,
132}
133impl From<wgt::QueryType> for SimplifiedQueryType {
134    fn from(q: wgt::QueryType) -> Self {
135        match q {
136            wgt::QueryType::Occlusion => SimplifiedQueryType::Occlusion,
137            wgt::QueryType::Timestamp => SimplifiedQueryType::Timestamp,
138            wgt::QueryType::PipelineStatistics(..) => SimplifiedQueryType::PipelineStatistics,
139        }
140    }
141}
142
143/// Error encountered when dealing with queries
144#[derive(Clone, Debug, Error)]
145#[non_exhaustive]
146pub enum QueryError {
147    #[error(transparent)]
148    Device(#[from] DeviceError),
149    #[error(transparent)]
150    EncoderState(#[from] EncoderStateError),
151    #[error(transparent)]
152    MissingFeature(#[from] MissingFeatures),
153    #[error("Error encountered while trying to use queries")]
154    Use(#[from] QueryUseError),
155    #[error("Error encountered while trying to resolve a query")]
156    Resolve(#[from] ResolveError),
157    #[error(transparent)]
158    DestroyedResource(#[from] DestroyedResourceError),
159    #[error(transparent)]
160    InvalidResource(#[from] InvalidResourceError),
161}
162
163impl WebGpuError for QueryError {
164    fn webgpu_error_type(&self) -> ErrorType {
165        match self {
166            Self::EncoderState(e) => e.webgpu_error_type(),
167            Self::Use(e) => e.webgpu_error_type(),
168            Self::Resolve(e) => e.webgpu_error_type(),
169            Self::InvalidResource(e) => e.webgpu_error_type(),
170            Self::Device(e) => e.webgpu_error_type(),
171            Self::MissingFeature(e) => e.webgpu_error_type(),
172            Self::DestroyedResource(e) => e.webgpu_error_type(),
173        }
174    }
175}
176
177/// Error encountered while trying to use queries
178#[derive(Clone, Debug, Error)]
179#[non_exhaustive]
180pub enum QueryUseError {
181    #[error(transparent)]
182    Device(#[from] DeviceError),
183    #[error("Query {query_index} is out of bounds for a query set of size {query_set_size}")]
184    OutOfBounds {
185        query_index: u32,
186        query_set_size: u32,
187    },
188    #[error("Query {query_index} has already been used within the same renderpass. Queries must only be used once per renderpass")]
189    UsedTwiceInsideRenderpass { query_index: u32 },
190    #[error("Query {new_query_index} was started while query {active_query_index} was already active. No more than one statistic or occlusion query may be active at once")]
191    AlreadyStarted {
192        active_query_index: u32,
193        new_query_index: u32,
194    },
195    #[error("Query was stopped while there was no active query")]
196    AlreadyStopped,
197    #[error("A query of type {query_type:?} was started using a query set of type {set_type:?}")]
198    IncompatibleType {
199        set_type: SimplifiedQueryType,
200        query_type: SimplifiedQueryType,
201    },
202    #[error("A query of type {query_type:?} was not ended before the encoder was finished")]
203    MissingEnd { query_type: SimplifiedQueryType },
204    #[error(transparent)]
205    DestroyedResource(#[from] DestroyedResourceError),
206}
207
208impl WebGpuError for QueryUseError {
209    fn webgpu_error_type(&self) -> ErrorType {
210        match self {
211            Self::Device(e) => e.webgpu_error_type(),
212            Self::DestroyedResource(e) => e.webgpu_error_type(),
213            Self::OutOfBounds { .. }
214            | Self::UsedTwiceInsideRenderpass { .. }
215            | Self::AlreadyStarted { .. }
216            | Self::AlreadyStopped
217            | Self::IncompatibleType { .. }
218            | Self::MissingEnd { .. } => ErrorType::Validation,
219        }
220    }
221}
222
223/// Error encountered while trying to resolve a query.
224#[derive(Clone, Debug, Error)]
225#[non_exhaustive]
226pub enum ResolveError {
227    #[error(transparent)]
228    MissingBufferUsage(#[from] MissingBufferUsageError),
229    #[error("Resolve buffer offset has to be aligned to `QUERY_RESOLVE_BUFFER_ALIGNMENT")]
230    BufferOffsetAlignment,
231    #[error("Resolving queries {start_query}..{end_query} would overrun the query set of size {query_set_size}")]
232    QueryOverrun {
233        start_query: u32,
234        end_query: u64,
235        query_set_size: u32,
236    },
237    #[error("Resolving queries {start_query}..{end_query} ({stride} byte queries) will end up overrunning the bounds of the destination buffer of size {buffer_size} using offsets {buffer_start_offset}..(<start> + {bytes_used})")]
238    BufferOverrun {
239        start_query: u32,
240        end_query: u32,
241        stride: u32,
242        buffer_size: BufferAddress,
243        buffer_start_offset: BufferAddress,
244        bytes_used: BufferAddress,
245    },
246}
247
248impl WebGpuError for ResolveError {
249    fn webgpu_error_type(&self) -> ErrorType {
250        match self {
251            Self::MissingBufferUsage(e) => e.webgpu_error_type(),
252            Self::BufferOffsetAlignment
253            | Self::QueryOverrun { .. }
254            | Self::BufferOverrun { .. } => ErrorType::Validation,
255        }
256    }
257}
258
259impl QuerySet {
260    pub(crate) fn validate_query(
261        self: &Arc<Self>,
262        query_type: SimplifiedQueryType,
263        query_index: u32,
264        reset_state: Option<&mut QueryResetMap>,
265    ) -> Result<(), QueryUseError> {
266        // NOTE: Further code assumes the index is good, so do this first.
267        if query_index >= self.desc.count {
268            return Err(QueryUseError::OutOfBounds {
269                query_index,
270                query_set_size: self.desc.count,
271            });
272        }
273
274        // We need to defer our resets because we are in a renderpass,
275        // add the usage to the reset map.
276        if let Some(reset) = reset_state {
277            let used = reset.use_query_set(self, query_index);
278            if used {
279                return Err(QueryUseError::UsedTwiceInsideRenderpass { query_index });
280            }
281        }
282
283        let simple_set_type = SimplifiedQueryType::from(self.desc.ty);
284        if simple_set_type != query_type {
285            return Err(QueryUseError::IncompatibleType {
286                query_type,
287                set_type: simple_set_type,
288            });
289        }
290
291        Ok(())
292    }
293
294    pub(super) fn validate_and_write_timestamp(
295        self: &Arc<Self>,
296        raw_encoder: &mut dyn hal::DynCommandEncoder,
297        query_index: u32,
298        reset_state: Option<&mut QueryResetMap>,
299        snatch_guard: &SnatchGuard<'_>,
300        query_set_writes: &mut QuerySetWrites,
301    ) -> Result<(), QueryUseError> {
302        let needs_reset = reset_state.is_none();
303        self.validate_query(SimplifiedQueryType::Timestamp, query_index, reset_state)?;
304
305        unsafe {
306            // If we don't have a reset state tracker which can defer resets, we must reset now.
307            if needs_reset {
308                raw_encoder
309                    .reset_queries(self.try_raw(snatch_guard)?, query_index..(query_index + 1));
310            }
311            raw_encoder.write_timestamp(self.try_raw(snatch_guard)?, query_index);
312        }
313
314        record_query_write(query_set_writes, self, query_index);
315        Ok(())
316    }
317}
318
319pub(super) fn validate_and_begin_occlusion_query(
320    query_set: Arc<QuerySet>,
321    raw_encoder: &mut dyn hal::DynCommandEncoder,
322    tracker: &mut QuerySetTracker,
323    query_index: u32,
324    reset_state: Option<&mut QueryResetMap>,
325    active_query: &mut Option<(Arc<QuerySet>, u32)>,
326    snatch_guard: &SnatchGuard<'_>,
327) -> Result<(), QueryUseError> {
328    let needs_reset = reset_state.is_none();
329    query_set.validate_query(SimplifiedQueryType::Occlusion, query_index, reset_state)?;
330
331    tracker.insert_single(query_set.clone());
332
333    if let Some((_old, old_idx)) = active_query.take() {
334        return Err(QueryUseError::AlreadyStarted {
335            active_query_index: old_idx,
336            new_query_index: query_index,
337        });
338    }
339    let (query_set, _) = &active_query.insert((query_set, query_index));
340
341    unsafe {
342        // If we don't have a reset state tracker which can defer resets, we must reset now.
343        if needs_reset {
344            raw_encoder.reset_queries(
345                query_set.try_raw(snatch_guard)?,
346                query_index..(query_index + 1),
347            );
348        }
349        raw_encoder.begin_query(query_set.try_raw(snatch_guard)?, query_index);
350    }
351
352    Ok(())
353}
354
355pub(super) fn end_occlusion_query(
356    raw_encoder: &mut dyn hal::DynCommandEncoder,
357    active_query: &mut Option<(Arc<QuerySet>, u32)>,
358    snatch_guard: &SnatchGuard<'_>,
359    query_set_writes: &mut QuerySetWrites,
360) -> Result<(), QueryUseError> {
361    if let Some((query_set, query_index)) = active_query.take() {
362        unsafe { raw_encoder.end_query(query_set.try_raw(snatch_guard)?, query_index) };
363        record_query_write(query_set_writes, &query_set, query_index);
364        Ok(())
365    } else {
366        Err(QueryUseError::AlreadyStopped)
367    }
368}
369
370pub(super) fn validate_and_begin_pipeline_statistics_query(
371    query_set: Arc<QuerySet>,
372    raw_encoder: &mut dyn hal::DynCommandEncoder,
373    tracker: &mut QuerySetTracker,
374    device: &Arc<Device>,
375    query_index: u32,
376    reset_state: Option<&mut QueryResetMap>,
377    active_query: &mut Option<(Arc<QuerySet>, u32)>,
378    snatch_guard: &SnatchGuard<'_>,
379) -> Result<(), QueryUseError> {
380    query_set.same_device(device)?;
381
382    let needs_reset = reset_state.is_none();
383    query_set.validate_query(
384        SimplifiedQueryType::PipelineStatistics,
385        query_index,
386        reset_state,
387    )?;
388
389    tracker.insert_single(query_set.clone());
390
391    if let Some((_old, old_idx)) = active_query.take() {
392        return Err(QueryUseError::AlreadyStarted {
393            active_query_index: old_idx,
394            new_query_index: query_index,
395        });
396    }
397    let (query_set, _) = &active_query.insert((query_set, query_index));
398
399    unsafe {
400        // If we don't have a reset state tracker which can defer resets, we must reset now.
401        if needs_reset {
402            raw_encoder.reset_queries(
403                query_set.try_raw(snatch_guard)?,
404                query_index..(query_index + 1),
405            );
406        }
407        raw_encoder.begin_query(query_set.try_raw(snatch_guard)?, query_index);
408    }
409
410    Ok(())
411}
412
413pub(super) fn end_pipeline_statistics_query(
414    raw_encoder: &mut dyn hal::DynCommandEncoder,
415    active_query: &mut Option<(Arc<QuerySet>, u32)>,
416    snatch_guard: &SnatchGuard<'_>,
417    query_set_writes: &mut QuerySetWrites,
418) -> Result<(), QueryUseError> {
419    if let Some((query_set, query_index)) = active_query.take() {
420        unsafe { raw_encoder.end_query(query_set.try_raw(snatch_guard)?, query_index) };
421        record_query_write(query_set_writes, &query_set, query_index);
422        Ok(())
423    } else {
424        Err(QueryUseError::AlreadyStopped)
425    }
426}
427
428impl super::CommandEncoder {
429    fn write_timestamp_inner(
430        self: &Arc<Self>,
431        query_set: Arc<QuerySet>,
432        query_index: u32,
433    ) -> Result<(), EncoderStateError> {
434        let mut cmd_buf_data = self.data.lock();
435
436        cmd_buf_data.push_with(|| -> Result<_, QueryError> {
437            query_set.check_is_valid()?;
438            Ok(ArcCommand::WriteTimestamp {
439                query_set,
440                query_index,
441            })
442        })
443    }
444
445    pub fn write_timestamp(self: &Arc<Self>, query_set: Arc<QuerySet>, query_index: u32) {
446        if let Err(err) = self.write_timestamp_inner(query_set, query_index) {
447            self.device
448                .handle_error(err, None, "CommandEncoder::write_timestamp");
449        }
450    }
451
452    fn resolve_query_set_inner(
453        self: &Arc<Self>,
454        query_set: Arc<QuerySet>,
455        start_query: u32,
456        query_count: u32,
457        destination: Arc<Buffer>,
458        destination_offset: BufferAddress,
459    ) -> Result<(), EncoderStateError> {
460        let mut cmd_buf_data = self.data.lock();
461
462        cmd_buf_data.push_with(|| -> Result<_, QueryError> {
463            query_set.check_is_valid()?;
464            destination.check_is_valid()?;
465            Ok(ArcCommand::ResolveQuerySet {
466                query_set,
467                start_query,
468                query_count,
469                destination,
470                destination_offset,
471            })
472        })
473    }
474
475    pub fn resolve_query_set(
476        self: &Arc<Self>,
477        query_set: Arc<QuerySet>,
478        start_query: u32,
479        query_count: u32,
480        destination: Arc<Buffer>,
481        destination_offset: BufferAddress,
482    ) {
483        if let Err(err) = self.resolve_query_set_inner(
484            query_set,
485            start_query,
486            query_count,
487            destination,
488            destination_offset,
489        ) {
490            self.device
491                .handle_error(err, Some(self.label()), "CommandEncoder::resolve_query_set");
492        }
493    }
494}
495
496pub(super) fn write_timestamp(
497    state: &mut EncodingState,
498    query_set: Arc<QuerySet>,
499    query_index: u32,
500) -> Result<(), QueryError> {
501    state
502        .device
503        .require_features(wgt::Features::TIMESTAMP_QUERY_INSIDE_ENCODERS)?;
504
505    query_set.same_device(state.device)?;
506
507    query_set.validate_and_write_timestamp(
508        state.raw_encoder,
509        query_index,
510        None,
511        state.snatch_guard,
512        state.query_set_writes,
513    )?;
514
515    state.tracker.query_sets.insert_single(query_set);
516
517    Ok(())
518}
519
520pub(super) fn resolve_query_set(
521    state: &mut EncodingState<'_, '_, InnerCommandEncoder>,
522    query_set: Arc<QuerySet>,
523    start_query: u32,
524    query_count: u32,
525    dst_buffer: Arc<Buffer>,
526    destination_offset: BufferAddress,
527) -> Result<(), QueryError> {
528    if !destination_offset.is_multiple_of(wgt::QUERY_RESOLVE_BUFFER_ALIGNMENT) {
529        return Err(QueryError::Resolve(ResolveError::BufferOffsetAlignment));
530    }
531
532    query_set.same_device(state.device)?;
533    dst_buffer.same_device(state.device)?;
534
535    dst_buffer.check_destroyed(state.snatch_guard)?;
536
537    let dst_pending = state
538        .tracker
539        .buffers
540        .set_single(&dst_buffer, wgt::BufferUses::COPY_DST);
541    let dst_barrier = dst_pending.map(|pending| pending.into_hal(&dst_buffer, state.snatch_guard));
542
543    dst_buffer
544        .check_usage(wgt::BufferUsages::QUERY_RESOLVE)
545        .map_err(ResolveError::MissingBufferUsage)?;
546
547    let end_query = u64::from(start_query)
548        .checked_add(u64::from(query_count))
549        .expect("`u64` overflow from adding two `u32`s, should be unreachable");
550    if end_query > u64::from(query_set.desc.count) {
551        return Err(ResolveError::QueryOverrun {
552            start_query,
553            end_query,
554            query_set_size: query_set.desc.count,
555        }
556        .into());
557    }
558    let end_query =
559        u32::try_from(end_query).expect("`u32` overflow for `end_query`, which should be `u32`");
560
561    let elements_per_query = match query_set.desc.ty {
562        wgt::QueryType::Occlusion => 1,
563        wgt::QueryType::PipelineStatistics(ps) => ps.bits().count_ones(),
564        wgt::QueryType::Timestamp => 1,
565    };
566    let stride = elements_per_query * wgt::QUERY_SIZE;
567    let bytes_used: BufferAddress = u64::from(stride)
568        .checked_mul(u64::from(query_count))
569        .expect("`stride` * `query_count` overflowed `u32`, should be unreachable");
570
571    let buffer_start_offset = destination_offset;
572    let buffer_end_offset = buffer_start_offset
573        .checked_add(bytes_used)
574        .filter(|buffer_end_offset| *buffer_end_offset <= dst_buffer.size)
575        .ok_or(ResolveError::BufferOverrun {
576            start_query,
577            end_query,
578            stride,
579            buffer_size: dst_buffer.size,
580            buffer_start_offset,
581            bytes_used,
582        })?;
583
584    let query_set = state.tracker.query_sets.insert_single(query_set);
585
586    state
587        .buffer_memory_init_actions
588        .extend(dst_buffer.initialization_status.read().create_action(
589            &dst_buffer,
590            buffer_start_offset..buffer_end_offset,
591            MemoryInitKind::ImplicitlyInitialized,
592        ));
593
594    let raw_encoder = state.raw_encoder.open_if_closed()?;
595    let raw_dst_buffer = dst_buffer.try_raw(state.snatch_guard)?;
596    unsafe {
597        raw_encoder.transition_buffers(dst_barrier.as_slice());
598    }
599
600    // Check if all slots in the range have been written within this encoder.
601    // If so we can emit `copy_query_results` directly.
602    // Otherwise defer to submit time where we have knowledge of
603    // the query set initialization state.
604    let query_set_writes = state.query_set_writes.get(&query_set.tracker_index());
605    let all_written =
606        query_set_writes.is_some_and(|slots| (start_query..end_query).all(|i| slots[i as usize]));
607
608    if all_written {
609        unsafe {
610            raw_encoder.copy_query_results(
611                query_set.try_raw(state.snatch_guard)?,
612                start_query..end_query,
613                raw_dst_buffer,
614                destination_offset,
615                wgt::BufferSize::new_unchecked(stride as u64),
616            );
617        }
618    } else {
619        state.raw_encoder.close_if_open()?;
620        let insertion_point = state.raw_encoder.list.len();
621
622        state
623            .deferred_query_set_resolves
624            .push(DeferredQuerySetResolve {
625                query_set: query_set.clone(),
626                query_set_writes: query_set_writes.cloned(),
627                start_query,
628                end_query,
629                dst_buffer: dst_buffer.clone(),
630                destination_offset,
631                stride: stride as u64,
632                insertion_point,
633            });
634    }
635
636    if matches!(query_set.desc.ty, wgt::QueryType::Timestamp) {
637        let raw_encoder = state.raw_encoder.open_if_closed()?;
638
639        // Timestamp normalization is only needed for timestamps.
640        state.device.timestamp_normalizer.get().unwrap().normalize(
641            state.snatch_guard,
642            raw_encoder,
643            &mut state.tracker.buffers,
644            dst_buffer
645                .timestamp_normalization_bind_group
646                .get(state.snatch_guard)
647                .unwrap(),
648            &dst_buffer,
649            destination_offset,
650            query_count,
651        );
652    }
653
654    Ok(())
655}