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