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#[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 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 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 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 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 pub const fn contains(self, other: Self) -> bool {
141 other.0 & !self.0 == 0
142 }
143
144 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}