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 pub(crate) stride: u64,
33 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 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 (Some(..), true) => {}
106 (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 (None, true) => {
116 run_start = Some(idx as u32);
117 }
118 (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#[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#[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#[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 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 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 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 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 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 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 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}