1use crate::maybe_pqp_cg_ssa as rustc_codegen_ssa;
3
4use super::CodegenCx;
5use crate::abi::ConvSpirvType;
6use crate::builder_spirv::{SpirvConst, SpirvValue, SpirvValueExt, SpirvValueKind};
7use crate::spirv_type::SpirvType;
8use itertools::Itertools as _;
9use rspirv::spirv::Word;
10use rustc_abi::{self as abi, AddressSpace, Float, HasDataLayout, Integer, Primitive, Size};
11use rustc_codegen_ssa::traits::{
12 ConstCodegenMethods, MiscCodegenMethods, PacMetadata, StaticCodegenMethods,
13};
14use rustc_middle::mir::interpret::{AllocError, ConstAllocation, GlobalAlloc, Scalar, alloc_range};
15use rustc_middle::ty::layout::LayoutOf;
16use rustc_span::{DUMMY_SP, Span};
17
18impl<'tcx> CodegenCx<'tcx> {
19 pub fn def_constant(&self, ty: Word, val: SpirvConst<'_, 'tcx>) -> SpirvValue {
20 self.builder.def_constant_cx(ty, val, self)
21 }
22
23 pub fn constant_u8(&self, span: Span, val: u8) -> SpirvValue {
24 self.constant_int_from_native_unsigned(span, val)
25 }
26
27 pub fn constant_i8(&self, span: Span, val: i8) -> SpirvValue {
28 self.constant_int_from_native_signed(span, val)
29 }
30
31 pub fn constant_i16(&self, span: Span, val: i16) -> SpirvValue {
32 self.constant_int_from_native_signed(span, val)
33 }
34
35 pub fn constant_u16(&self, span: Span, val: u16) -> SpirvValue {
36 self.constant_int_from_native_unsigned(span, val)
37 }
38
39 pub fn constant_i32(&self, span: Span, val: i32) -> SpirvValue {
40 self.constant_int_from_native_signed(span, val)
41 }
42
43 pub fn constant_u32(&self, span: Span, val: u32) -> SpirvValue {
44 self.constant_int_from_native_unsigned(span, val)
45 }
46
47 pub fn constant_i64(&self, span: Span, val: i64) -> SpirvValue {
48 self.constant_int_from_native_signed(span, val)
49 }
50
51 pub fn constant_u64(&self, span: Span, val: u64) -> SpirvValue {
52 self.constant_int_from_native_unsigned(span, val)
53 }
54
55 pub fn constant_u128(&self, span: Span, val: u128) -> SpirvValue {
56 self.constant_int_from_native_unsigned(span, val)
57 }
58
59 fn constant_int_from_native_unsigned(&self, span: Span, val: impl Into<u128>) -> SpirvValue {
60 let size = Size::from_bytes(std::mem::size_of_val(&val));
61 let ty = SpirvType::Integer(size.bits() as u32, false).def(span, self);
62 self.constant_int(ty, val.into())
63 }
64
65 fn constant_int_from_native_signed(&self, span: Span, val: impl Into<i128>) -> SpirvValue {
66 let size = Size::from_bytes(std::mem::size_of_val(&val));
67 let ty = SpirvType::Integer(size.bits() as u32, true).def(span, self);
68 self.constant_int(ty, val.into() as u128)
69 }
70
71 pub fn constant_int(&self, ty: Word, val: u128) -> SpirvValue {
72 self.def_constant(ty, SpirvConst::Scalar(val))
73 }
74
75 pub fn constant_f32(&self, span: Span, val: f32) -> SpirvValue {
76 let ty = SpirvType::Float(32).def(span, self);
77 self.def_constant(ty, SpirvConst::Scalar(val.to_bits().into()))
78 }
79
80 pub fn constant_f64(&self, span: Span, val: f64) -> SpirvValue {
81 let ty = SpirvType::Float(64).def(span, self);
82 self.def_constant(ty, SpirvConst::Scalar(val.to_bits().into()))
83 }
84
85 pub fn constant_float(&self, ty: Word, val: f64) -> SpirvValue {
86 match self.lookup_type(ty) {
87 SpirvType::Float(32) => {
89 self.def_constant(ty, SpirvConst::Scalar((val as f32).to_bits().into()))
90 }
91 SpirvType::Float(64) => self.def_constant(ty, SpirvConst::Scalar(val.to_bits().into())),
92 other => self.tcx.dcx().fatal(format!(
93 "constant_float does not support type {}",
94 other.debug(ty, self)
95 )),
96 }
97 }
98
99 pub fn constant_bool(&self, span: Span, val: bool) -> SpirvValue {
100 let ty = SpirvType::Bool.def(span, self);
101 self.def_constant(ty, SpirvConst::Scalar(val as u128))
102 }
103
104 pub fn constant_composite(&self, ty: Word, fields: impl Iterator<Item = Word>) -> SpirvValue {
105 self.def_constant(ty, SpirvConst::Composite(&fields.collect::<Vec<_>>()))
107 }
108
109 pub fn constant_null(&self, ty: Word) -> SpirvValue {
110 self.def_constant(ty, SpirvConst::Null)
111 }
112
113 pub fn undef(&self, ty: Word) -> SpirvValue {
114 self.def_constant(ty, SpirvConst::Undef)
115 }
116}
117
118impl ConstCodegenMethods for CodegenCx<'_> {
119 fn const_null(&self, t: Self::Type) -> Self::Value {
120 self.constant_null(t)
121 }
122 fn const_undef(&self, ty: Self::Type) -> Self::Value {
123 self.undef(ty)
124 }
125 fn const_poison(&self, ty: Self::Type) -> Self::Value {
126 self.const_undef(ty)
128 }
129 fn const_int(&self, t: Self::Type, i: i64) -> Self::Value {
130 self.constant_int(t, i as u128)
131 }
132 fn const_uint(&self, t: Self::Type, i: u64) -> Self::Value {
133 self.constant_int(t, i.into())
134 }
135 fn const_uint_big(&self, t: Self::Type, i: u128) -> Self::Value {
136 self.constant_int(t, i)
137 }
138 fn const_bool(&self, val: bool) -> Self::Value {
139 self.constant_bool(DUMMY_SP, val)
140 }
141 fn const_i8(&self, i: i8) -> Self::Value {
142 self.constant_i8(DUMMY_SP, i)
143 }
144 fn const_i16(&self, i: i16) -> Self::Value {
145 self.constant_i16(DUMMY_SP, i)
146 }
147 fn const_i32(&self, i: i32) -> Self::Value {
148 self.constant_i32(DUMMY_SP, i)
149 }
150 fn const_i64(&self, i: i64) -> Self::Value {
151 self.constant_i64(DUMMY_SP, i)
152 }
153 fn const_u8(&self, i: u8) -> Self::Value {
154 self.constant_u8(DUMMY_SP, i)
155 }
156 fn const_u32(&self, i: u32) -> Self::Value {
157 self.constant_u32(DUMMY_SP, i)
158 }
159 fn const_u64(&self, i: u64) -> Self::Value {
160 self.constant_u64(DUMMY_SP, i)
161 }
162 fn const_u128(&self, i: u128) -> Self::Value {
163 let ty = SpirvType::Integer(128, false).def(DUMMY_SP, self);
164 self.const_uint_big(ty, i)
165 }
166 fn const_usize(&self, i: u64) -> Self::Value {
167 let ptr_size = self.tcx.data_layout.pointer_size().bits() as u32;
168 let t = SpirvType::Integer(ptr_size, false).def(DUMMY_SP, self);
169 self.constant_int(t, i.into())
170 }
171 fn const_real(&self, t: Self::Type, val: f64) -> Self::Value {
172 self.constant_float(t, val)
173 }
174
175 fn const_str(&self, s: &str) -> (Self::Value, Self::Value) {
176 let len = s.len();
177 let str_ty = self
178 .layout_of(self.tcx.types.str_)
179 .spirv_type(DUMMY_SP, self);
180 (
181 self.def_constant(
182 self.type_ptr_to(str_ty),
183 SpirvConst::PtrTo {
184 pointee: self
185 .constant_composite(
186 str_ty,
187 s.bytes().map(|b| self.const_u8(b).def_cx(self)),
188 )
189 .def_cx(self),
190 },
191 ),
192 self.const_usize(len as u64),
193 )
194 }
195 fn const_struct(&self, elts: &[Self::Value], _packed: bool) -> Self::Value {
196 let field_types = elts.iter().map(|f| f.ty).collect::<Vec<_>>();
199 let (field_offsets, size, align) = crate::abi::auto_struct_layout(self, &field_types);
200 let struct_ty = SpirvType::Adt {
201 def_id: None,
202 size,
203 align,
204 field_types: &field_types,
205 field_offsets: &field_offsets,
206 field_names: None,
207 }
208 .def(DUMMY_SP, self);
209 self.constant_composite(struct_ty, elts.iter().map(|f| f.def_cx(self)))
210 }
211 fn const_vector(&self, elts: &[Self::Value]) -> Self::Value {
212 let vector_ty = SpirvType::simd_vector(
213 self,
214 DUMMY_SP,
215 self.lookup_type(elts[0].ty),
216 elts.len() as u32,
217 )
218 .def(DUMMY_SP, self);
219 self.constant_composite(vector_ty, elts.iter().map(|elt| elt.def_cx(self)))
220 }
221
222 fn const_to_opt_uint(&self, v: Self::Value) -> Option<u64> {
223 self.builder.lookup_const_scalar(v)?.try_into().ok()
224 }
225 fn const_to_opt_u128(&self, v: Self::Value, _sign_ext: bool) -> Option<u128> {
228 self.builder.lookup_const_scalar(v)
229 }
230
231 fn scalar_to_backend_with_pac(
232 &self,
233 cv: Scalar,
234 layout: rustc_abi::Scalar,
235 ty: Self::Type,
236 _pac: Option<PacMetadata>,
237 ) -> Self::Value {
238 self.scalar_to_backend(cv, layout, ty)
239 }
240
241 fn scalar_to_backend(
242 &self,
243 scalar: Scalar,
244 layout: abi::Scalar,
245 ty: Self::Type,
246 ) -> Self::Value {
247 match scalar {
248 Scalar::Int(int) => {
249 assert_eq!(int.size(), layout.primitive().size(self));
250 let data = int.to_uint(int.size());
251
252 if let Primitive::Pointer(_) = layout.primitive() {
253 if data == 0 {
254 self.constant_null(ty)
255 } else {
256 let result = self.undef(ty);
257 self.zombie_no_span(
258 result.def_cx(self),
259 "pointer has non-null integer address",
260 );
261 result
262 }
263 } else {
264 self.def_constant(ty, SpirvConst::Scalar(data))
265 }
266 }
267 Scalar::Ptr(ptr, _) => {
268 let (prov, offset) = ptr.prov_and_relative_offset();
269 let alloc_id = prov.alloc_id();
270 let (base_addr, _base_addr_space) = match self.tcx.global_alloc(alloc_id) {
271 GlobalAlloc::Memory(alloc) => {
272 match self.lookup_type(ty) {
273 SpirvType::Pointer { .. } => {}
274 other => self.tcx.dcx().fatal(format!(
275 "GlobalAlloc::Memory type not implemented: {}",
276 other.debug(ty, self)
277 )),
278 }
279 let value = self.static_addr_of(alloc, None);
280 (value, AddressSpace::ZERO)
281 }
282 GlobalAlloc::Function { instance } => (
283 self.get_fn_addr(instance, None),
284 self.data_layout().instruction_address_space,
285 ),
286 GlobalAlloc::VTable(vty, dyn_ty) => {
287 let alloc = self
288 .tcx
289 .global_alloc(self.tcx.vtable_allocation((
290 vty,
291 dyn_ty.principal().map(|principal| {
292 self.tcx.instantiate_bound_regions_with_erased(principal)
293 }),
294 )))
295 .unwrap_memory();
296 match self.lookup_type(ty) {
297 SpirvType::Pointer { .. } => {}
298 other => self.tcx.dcx().fatal(format!(
299 "GlobalAlloc::VTable type not implemented: {}",
300 other.debug(ty, self)
301 )),
302 }
303 let value = self.static_addr_of(alloc, None);
304 (value, AddressSpace::ZERO)
305 }
306 GlobalAlloc::Static(def_id) => {
307 assert!(self.tcx.is_static(def_id));
308 assert!(!self.tcx.is_thread_local_static(def_id));
309 (self.get_static(def_id), AddressSpace::ZERO)
310 }
311 GlobalAlloc::TypeId { .. } => {
312 return if offset.bytes() == 0 {
313 self.constant_null(ty)
314 } else {
315 let result = self.undef(ty);
316 self.zombie_no_span(
317 result.def_cx(self),
318 "pointer has non-null integer address",
319 );
320 result
321 };
322 }
323 };
324 self.const_bitcast(self.const_ptr_byte_offset(base_addr, offset), ty)
325 }
326 }
327 }
328
329 fn const_ptr_byte_offset(&self, val: Self::Value, offset: Size) -> Self::Value {
330 if offset == Size::ZERO {
331 val
332 } else {
333 let result = val;
336 self.zombie_no_span(result.def_cx(self), "const_ptr_byte_offset");
337 result
338 }
339 }
340}
341
342impl<'tcx> CodegenCx<'tcx> {
343 pub(crate) fn const_data_from_alloc(&self, alloc: ConstAllocation<'_>) -> SpirvValue {
348 let alloc = self.tcx.lift(alloc);
352
353 let void_type = SpirvType::Void.def(DUMMY_SP, self);
354 self.def_constant(void_type, SpirvConst::ConstDataFromAlloc(alloc))
355 }
356
357 pub fn const_bitcast(&self, val: SpirvValue, ty: Word) -> SpirvValue {
358 if let SpirvValueKind::IllegalConst(_) = val.kind
361 && let Some(SpirvConst::PtrTo { pointee }) = self.builder.lookup_const(val)
362 && let Some(SpirvConst::ConstDataFromAlloc(alloc)) =
363 self.builder.lookup_const_by_id(pointee)
364 && let SpirvType::Pointer { pointee } = self.lookup_type(ty)
365 && let Some(init) = self.try_read_from_const_alloc(alloc, pointee)
366 {
367 return self.def_constant(
368 ty,
369 SpirvConst::PtrTo {
370 pointee: init.def_cx(self),
371 },
372 );
373 }
374
375 if val.ty == ty {
376 val
377 } else {
378 let result = val.def_cx(self).with_type(ty);
381 self.zombie_no_span(result.def_cx(self), "const_bitcast");
382 result
383 }
384 }
385
386 pub fn primitive_to_scalar(&self, value: Primitive) -> abi::Scalar {
389 let bits = value.size(self.data_layout()).bits();
390 assert!(bits <= 128);
391 abi::Scalar::Initialized {
392 value,
393 valid_range: abi::WrappingRange {
394 start: 0,
395 end: (!0 >> (128 - bits)),
396 },
397 }
398 }
399
400 pub fn try_read_from_const_alloc(
405 &self,
406 alloc: ConstAllocation<'tcx>,
407 ty: Word,
408 ) -> Option<SpirvValue> {
409 let (result, read_size) = self.read_from_const_alloc_at(alloc, ty, Size::ZERO);
410 (read_size == alloc.inner().size()).then_some(result)
411 }
412
413 #[tracing::instrument(level = "trace", skip(self), fields(ty = ?self.debug_type(ty), offset))]
418 fn read_from_const_alloc_at(
419 &self,
420 alloc: ConstAllocation<'tcx>,
421 ty: Word,
422 offset: Size,
423 ) -> (SpirvValue, Size) {
424 let ty_def = self.lookup_type(ty);
425 match ty_def {
426 SpirvType::Bool
427 | SpirvType::Integer(..)
428 | SpirvType::Float(_)
429 | SpirvType::Pointer { .. } => {
430 let size = ty_def.sizeof(self).unwrap();
431 let primitive = match ty_def {
432 SpirvType::Bool => Primitive::Int(Integer::fit_unsigned(0), false),
433 SpirvType::Integer(int_size, int_signedness) => Primitive::Int(
434 match int_size {
435 8 => Integer::I8,
436 16 => Integer::I16,
437 32 => Integer::I32,
438 64 => Integer::I64,
439 128 => Integer::I128,
440 other => {
441 self.tcx
442 .dcx()
443 .fatal(format!("invalid size for integer: {other}"));
444 }
445 },
446 int_signedness,
447 ),
448 SpirvType::Float(float_size) => Primitive::Float(match float_size {
449 16 => Float::F16,
450 32 => Float::F32,
451 64 => Float::F64,
452 128 => Float::F128,
453 other => {
454 self.tcx
455 .dcx()
456 .fatal(format!("invalid size for float: {other}"));
457 }
458 }),
459 SpirvType::Pointer { .. } => Primitive::Pointer(AddressSpace::ZERO),
460 _ => unreachable!(),
461 };
462
463 let range = alloc_range(offset, size);
464 let read_provenance = matches!(primitive, Primitive::Pointer(_));
465
466 let mut primitive = primitive;
467 let mut read_result = alloc.inner().read_scalar(self, range, read_provenance);
468
469 if read_result.is_err()
473 && !read_provenance
474 && let read_ptr_result @ Ok(Scalar::Ptr(ptr, _)) = alloc
475 .inner()
476 .read_scalar(self, range, true)
477 {
478 let (prov, _offset) = ptr.prov_and_relative_offset();
479 primitive = Primitive::Pointer(
480 self.tcx.global_alloc(prov.alloc_id()).address_space(self),
481 );
482 read_result = read_ptr_result;
483 }
484
485 let scalar_or_zombie = match read_result {
486 Ok(scalar) => {
487 Ok(self.scalar_to_backend(scalar, self.primitive_to_scalar(primitive), ty))
488 }
489
490 Err(err) => match err {
493 AllocError::InvalidUninitBytes(_) => {
497 let uninit_range = alloc
498 .inner()
499 .init_mask()
500 .is_range_initialized(range)
501 .unwrap_err();
502 let uninit_size = {
503 let [start, end] = [uninit_range.start, uninit_range.end()]
504 .map(|x| x.clamp(range.start, range.end()));
505 end - start
506 };
507 if uninit_size == size {
508 Ok(self.undef(ty))
509 } else {
510 Err(format!(
511 "overlaps {} uninitialized bytes",
512 uninit_size.bytes()
513 ))
514 }
515 }
516 AllocError::ReadPointerAsInt(_) => Err("overlaps pointer bytes".into()),
517 AllocError::ReadPartialPointer(_) => {
518 Err("partially overlaps another pointer".into())
519 }
520
521 AllocError::ScalarSizeMismatch(_) => {
524 Err(format!("unrecognized `AllocError::{err:?}`"))
525 }
526 },
527 };
528 let result = scalar_or_zombie.unwrap_or_else(|reason| {
529 let result = self.undef(ty);
530 self.zombie_no_span(
531 result.def_cx(self),
532 &format!("unsupported `{}` constant: {reason}", self.debug_type(ty),),
533 );
534 result
535 });
536 (result, size)
537 }
538 SpirvType::Adt {
539 field_types,
540 field_offsets,
541 ..
542 } => {
543 let mut tail_read_range = ..Size::ZERO;
546 let result = self.constant_composite(
547 ty,
548 field_types
549 .iter()
550 .zip_eq(field_offsets.iter())
551 .map(|(&f_ty, &f_offset)| {
552 let (f, f_size) =
553 self.read_from_const_alloc_at(alloc, f_ty, offset + f_offset);
554 tail_read_range.end =
555 tail_read_range.end.max(offset + f_offset + f_size);
556 f.def_cx(self)
557 }),
558 );
559
560 let ty_size = ty_def.sizeof(self);
561
562 if let Some(ty_size) = ty_size
564 && let Some(tail_gap) = (ty_size.bytes())
565 .checked_sub(tail_read_range.end.align_to(ty_def.alignof(self)).bytes())
566 && tail_gap > 0
567 {
568 self.zombie_no_span(
569 result.def_cx(self),
570 &format!(
571 "undersized `{}` constant (at least {tail_gap} bytes may be missing)",
572 self.debug_type(ty)
573 ),
574 );
575 }
576
577 (result, ty_size.unwrap_or(tail_read_range.end))
578 }
579 SpirvType::Vector { element, .. }
580 | SpirvType::Matrix { element, .. }
581 | SpirvType::Array { element, .. }
582 | SpirvType::RuntimeArray { element } => {
583 let stride = self.lookup_type(element).sizeof(self).unwrap();
584
585 let count = match ty_def {
586 SpirvType::Vector { count, .. } | SpirvType::Matrix { count, .. } => {
587 u64::from(count)
588 }
589 SpirvType::Array { count, .. } => {
590 u64::try_from(self.builder.lookup_const_scalar(count).unwrap()).unwrap()
591 }
592 SpirvType::RuntimeArray { .. } => {
593 (alloc.inner().size() - offset).bytes() / stride.bytes()
594 }
595 _ => unreachable!(),
596 };
597
598 let result = self.constant_composite(
599 ty,
600 (0..count).map(|i| {
601 let (e, e_size) =
602 self.read_from_const_alloc_at(alloc, element, offset + i * stride);
603 assert_eq!(e_size, stride);
604 e.def_cx(self)
605 }),
606 );
607
608 let read_size = (count * stride).align_to(ty_def.alignof(self));
612
613 if let Some(ty_size) = ty_def.sizeof(self) {
614 assert_eq!(read_size, ty_size);
615 }
616
617 if let SpirvType::RuntimeArray { .. } = ty_def {
618 self.zombie_no_span(
623 result.def_cx(self),
624 &format!("unsupported unsized `{}` constant", self.debug_type(ty)),
625 );
626 }
627
628 (result, read_size)
629 }
630
631 SpirvType::Void
632 | SpirvType::Function { .. }
633 | SpirvType::Image { .. }
634 | SpirvType::Sampler
635 | SpirvType::SampledImage { .. }
636 | SpirvType::InterfaceBlock { .. }
637 | SpirvType::AccelerationStructureKhr
638 | SpirvType::RayQueryKhr
639 | SpirvType::CooperativeMatrixKhr { .. } => {
640 let result = self.undef(ty);
641 self.zombie_no_span(
642 result.def_cx(self),
643 &format!(
644 "cannot reinterpret Rust constant data as a `{}` value",
645 self.debug_type(ty)
646 ),
647 );
648 (result, ty_def.sizeof(self).unwrap_or(Size::ZERO))
649 }
650 }
651 }
652}