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 pub(crate) stride: u64,
35 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 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 (Some(..), true) => {}
108 (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 (None, true) => {
118 run_start = Some(idx as u32);
119 }
120 (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#[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#[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#[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 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 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 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 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 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 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 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}