naga/valid/
immediates.rs

1use core::{
2    fmt::{self, Debug},
3    ops,
4};
5
6#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
7#[error("Immediate size {0} overflows the bitmask")]
8pub struct ImmediateSlotsOverflowError(pub u64);
9
10/// A bitmask, tracking which 4-byte slots have been written via `set_immediates`.
11/// Bit N corresponds to bytes [N*4 .. N*4+4).
12#[derive(Clone, Copy, Default, PartialEq, Eq)]
13pub struct ImmediateSlots(u64);
14
15#[derive(Clone, Copy, Debug)]
16pub enum ImmediateUsage {
17    Valid { slots: ImmediateSlots, size: u32 },
18    // Used when a shader's immediate data type exceeds the maximum of 256 bytes and cannot
19    // be represented in the `ImmediateSlots` bitmask.
20    Invalid { size: u32 },
21}
22
23impl Default for ImmediateUsage {
24    fn default() -> Self {
25        ImmediateUsage::Valid {
26            slots: ImmediateSlots::default(),
27            size: 0,
28        }
29    }
30}
31
32impl ImmediateUsage {
33    pub const fn size(&self) -> u32 {
34        match *self {
35            ImmediateUsage::Valid { size, .. } => size,
36            ImmediateUsage::Invalid { size } => size,
37        }
38    }
39
40    pub fn from_type(
41        ty: &crate::TypeInner,
42        types: &crate::UniqueArena<crate::Type>,
43        gctx: crate::proc::GlobalCtx,
44    ) -> Self {
45        let size = ty.size(gctx);
46        ImmediateSlots::from_type(ty, types, gctx)
47            .map(|slots| Self::Valid { slots, size })
48            .unwrap_or(Self::Invalid { size })
49    }
50
51    pub fn merge(&self, other: &ImmediateUsage) -> Self {
52        let size = self.size().max(other.size());
53        match (*self, *other) {
54            (
55                ImmediateUsage::Valid { slots, .. },
56                ImmediateUsage::Valid {
57                    slots: other_slots, ..
58                },
59            ) => Self::Valid {
60                slots: slots | other_slots,
61                size,
62            },
63            _ => Self::Invalid { size },
64        }
65    }
66}
67
68impl ImmediateSlots {
69    pub const fn from_raw(raw: u64) -> Self {
70        Self(raw)
71    }
72
73    /// Compute the bitmask for a byte range [offset .. offset + size_bytes).
74    pub const fn from_range(
75        offset: u32,
76        size_bytes: u32,
77    ) -> Result<Self, ImmediateSlotsOverflowError> {
78        let Some(end) = offset.checked_add(size_bytes) else {
79            return Err(ImmediateSlotsOverflowError(
80                offset as u64 + size_bytes as u64,
81            ));
82        };
83        if end > u64::BITS * 4 {
84            return Err(ImmediateSlotsOverflowError(end as u64));
85        }
86        if size_bytes == 0 {
87            return Ok(Self(0));
88        }
89        let lo = offset / 4;
90        let hi = (offset + size_bytes).div_ceil(4);
91        Ok(Self(u64::MAX << lo & u64::MAX >> (64 - hi)))
92    }
93
94    /// Compute the slots occupied by a type,
95    /// excluding padding between matrix columns or struct members.
96    pub fn from_type(
97        ty: &crate::TypeInner,
98        types: &crate::UniqueArena<crate::Type>,
99        gctx: crate::proc::GlobalCtx,
100    ) -> Result<Self, ImmediateSlotsOverflowError> {
101        fn from_type_recursive(
102            ty: &crate::TypeInner,
103            offset: u32,
104            types: &crate::UniqueArena<crate::Type>,
105            gctx: crate::proc::GlobalCtx,
106        ) -> Result<ImmediateSlots, ImmediateSlotsOverflowError> {
107            // <https://www.w3.org/TR/WGSL/#accessible-bytes>
108            match *ty {
109                crate::TypeInner::Matrix {
110                    columns,
111                    rows,
112                    scalar,
113                } => {
114                    let mut slots = ImmediateSlots::default();
115                    let stride = crate::proc::Alignment::from(rows) * u32::from(scalar.width);
116                    for col in 0..u32::from(columns) {
117                        slots |= ImmediateSlots::from_range(
118                            offset + col * stride,
119                            u32::from(rows) * u32::from(scalar.width),
120                        )?;
121                    }
122                    Ok(slots)
123                }
124                crate::TypeInner::Struct { ref members, .. } => {
125                    let mut slots = ImmediateSlots::default();
126                    for member in members {
127                        let member_ty = &types[member.ty].inner;
128                        slots |=
129                            from_type_recursive(member_ty, offset + member.offset, types, gctx)?;
130                    }
131                    Ok(slots)
132                }
133                _ => ImmediateSlots::from_range(offset, ty.size(gctx)),
134            }
135        }
136        from_type_recursive(ty, 0, types, gctx)
137    }
138
139    /// Returns true if `self` contains all bits in `other`.
140    pub const fn contains(self, other: Self) -> bool {
141        other.0 & !self.0 == 0
142    }
143
144    /// Returns the bits in `self` that are not set in `other`.
145    pub const fn difference(self, other: Self) -> Self {
146        Self(self.0 & !other.0)
147    }
148}
149
150impl ops::BitOrAssign for ImmediateSlots {
151    fn bitor_assign(&mut self, rhs: Self) {
152        self.0 |= rhs.0;
153    }
154}
155
156impl ops::BitOr for ImmediateSlots {
157    type Output = Self;
158    fn bitor(self, rhs: Self) -> Self {
159        Self(self.0 | rhs.0)
160    }
161}
162
163impl fmt::Display for ImmediateSlots {
164    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
165        if self.0 == 0 {
166            return write!(f, "(none)");
167        }
168        let mut first = true;
169        let mut bit = 0u32;
170        while bit < 64 {
171            if self.0 & (1u64 << bit) != 0 {
172                let start = bit * 4;
173                while bit < 64 && self.0 & (1u64 << bit) != 0 {
174                    bit += 1;
175                }
176                let end = bit * 4;
177                if !first {
178                    write!(f, ", ")?;
179                }
180                write!(f, "{start}..{end}")?;
181                first = false;
182            } else {
183                bit += 1;
184            }
185        }
186        Ok(())
187    }
188}
189
190impl Debug for ImmediateSlots {
191    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
192        write!(f, "{self}")
193    }
194}
195
196#[cfg(test)]
197mod tests {
198    use crate::valid::ImmediateSlotsOverflowError;
199
200    use super::ImmediateSlots;
201
202    #[test]
203    fn range_single() {
204        assert_eq!(
205            ImmediateSlots::from_range(0, 4).unwrap(),
206            ImmediateSlots::from_raw(0b1)
207        );
208        assert_eq!(
209            ImmediateSlots::from_range(4, 4).unwrap(),
210            ImmediateSlots::from_raw(0b10)
211        );
212        assert_eq!(
213            ImmediateSlots::from_range(8, 4).unwrap(),
214            ImmediateSlots::from_raw(0b100)
215        );
216    }
217
218    #[test]
219    fn range_vec4() {
220        assert_eq!(
221            ImmediateSlots::from_range(0, 16).unwrap(),
222            ImmediateSlots::from_raw(0b1111)
223        );
224        assert_eq!(
225            ImmediateSlots::from_range(16, 16).unwrap(),
226            ImmediateSlots::from_raw(0b1111_0000)
227        );
228    }
229
230    #[test]
231    fn range_full_256() {
232        assert_eq!(
233            ImmediateSlots::from_range(0, 256).unwrap(),
234            ImmediateSlots::from_raw(u64::MAX)
235        );
236    }
237
238    #[test]
239    fn range_overflow() {
240        assert_eq!(
241            ImmediateSlots::from_range(0, 257),
242            Err(ImmediateSlotsOverflowError(257))
243        );
244    }
245
246    #[test]
247    fn from_type_overflow() {
248        let module = crate::front::wgsl::parse_str(
249            "struct S { \
250            e64: mat4x4<f32>, \
251            e128: mat4x4<f32>, \
252            e192: mat4x4<f32>, \
253            e256: mat4x4<f32>, \
254            e260: f32\
255            }",
256        )
257        .unwrap();
258        let struct_ty = (module.types.iter().map(|ty| ty.1))
259            .find(|ty| ty.name.as_deref() == Some("S"))
260            .unwrap();
261        let slots = ImmediateSlots::from_type(&struct_ty.inner, &module.types, module.to_ctx());
262        assert_eq!(slots, Err(ImmediateSlotsOverflowError(260)));
263    }
264
265    #[test]
266    fn from_type_excludes_struct_padding() {
267        let module = crate::front::wgsl::parse_str("struct S { a: f32, b: vec4<f32> }").unwrap();
268        let struct_ty = (module.types.iter().map(|ty| ty.1))
269            .find(|ty| ty.name.as_deref() == Some("S"))
270            .unwrap();
271        let slots =
272            ImmediateSlots::from_type(&struct_ty.inner, &module.types, module.to_ctx()).unwrap();
273        assert_eq!(slots, ImmediateSlots::from_raw(0b1111_0001));
274    }
275
276    #[test]
277    fn from_type_excludes_matrix_padding() {
278        let module = crate::front::wgsl::parse_str("struct S { mat: mat3x3<f32> }").unwrap();
279        let struct_ty = (module.types.iter().map(|ty| ty.1))
280            .find(|ty| ty.name.as_deref() == Some("S"))
281            .unwrap();
282        let slots =
283            ImmediateSlots::from_type(&struct_ty.inner, &module.types, module.to_ctx()).unwrap();
284        assert_eq!(slots, ImmediateSlots::from_raw(0b0111_0111_0111));
285
286        let module =
287            crate::front::wgsl::parse_str("struct S { f: f32, mat: mat2x2<f32> }").unwrap();
288        let struct_ty = (module.types.iter().map(|ty| ty.1))
289            .find(|ty| ty.name.as_deref() == Some("S"))
290            .unwrap();
291        let slots =
292            ImmediateSlots::from_type(&struct_ty.inner, &module.types, module.to_ctx()).unwrap();
293        assert_eq!(slots, ImmediateSlots::from_raw(0b11_11_01));
294    }
295
296    #[test]
297    fn range_unaligned() {
298        assert_eq!(
299            ImmediateSlots::from_range(0, 3).unwrap(),
300            ImmediateSlots::from_raw(0b1)
301        );
302        assert_eq!(
303            ImmediateSlots::from_range(0, 5).unwrap(),
304            ImmediateSlots::from_raw(0b11)
305        );
306    }
307
308    #[test]
309    fn contains() {
310        let required = ImmediateSlots::from_raw(0b1111_0001);
311        let mut set = ImmediateSlots::default();
312        assert!(!set.contains(required));
313        set |= ImmediateSlots::from_range(0, 4).unwrap();
314        assert!(!set.contains(required));
315        set |= ImmediateSlots::from_range(16, 16).unwrap();
316        assert!(set.contains(required));
317    }
318
319    #[test]
320    fn difference() {
321        let required = ImmediateSlots::from_raw(0b1111_0001);
322        let set = ImmediateSlots::from_range(0, 4).unwrap();
323        assert_eq!(
324            required.difference(set),
325            ImmediateSlots::from_raw(0b1111_0000)
326        );
327    }
328}