1use alloc::{boxed::Box, string::String, sync::Arc, vec::Vec};
2
3use miden_debug_types::{SourceManager, SourceSpan, Span, Spanned};
4use midenc_hir_type::{AddressSpace, Type, TypeRepr, TypeTemplate};
5
6use super::{
7 ConstantExpr, DocString, GlobalItemIndex, Ident, ItemIndex, Path, SymbolResolution,
8 SymbolResolutionError, Visibility, types,
9};
10
11pub(crate) const MAX_TYPE_EXPR_NESTING: usize = 256;
16
17pub trait TypeResolver<E> {
31 fn source_manager(&self) -> Arc<dyn SourceManager>;
32 fn resolve_local_failed(&self, err: SymbolResolutionError) -> E;
35 fn get_type(
42 &mut self,
43 context: SourceSpan,
44 gid: GlobalItemIndex,
45 ) -> Result<Option<TypeTemplate>, E>;
46 fn get_local_type(
48 &mut self,
49 context: SourceSpan,
50 id: ItemIndex,
51 ) -> Result<Option<TypeTemplate>, E>;
52 fn resolve_type_ref(&mut self, ty: Span<&Path>) -> Result<SymbolResolution, E>;
54 fn finalize(&mut self, context: SourceSpan, template: TypeTemplate) -> Result<Type, E>;
56 fn resolve(&mut self, ty: &TypeExpr) -> Result<Option<Type>, E> {
58 match ty.resolve_template(self)? {
59 Some(template) => self.finalize(ty.span(), template).map(Some),
60 None => Ok(None),
61 }
62 }
63}
64
65#[derive(Debug, Clone, PartialEq, Eq)]
70pub enum TypeDecl {
71 Alias(TypeAlias),
73 Enum(EnumType),
75}
76
77impl TypeDecl {
78 pub fn with_docs(self, docs: Option<Span<String>>) -> Self {
80 match self {
81 Self::Alias(ty) => Self::Alias(ty.with_docs(docs)),
82 Self::Enum(ty) => Self::Enum(ty.with_docs(docs)),
83 }
84 }
85
86 pub fn name(&self) -> &Ident {
88 match self {
89 Self::Alias(ty) => &ty.name,
90 Self::Enum(ty) => &ty.name,
91 }
92 }
93
94 pub const fn visibility(&self) -> Visibility {
96 match self {
97 Self::Alias(ty) => ty.visibility,
98 Self::Enum(ty) => ty.visibility,
99 }
100 }
101
102 pub fn docs(&self) -> Option<Span<&str>> {
104 match self {
105 Self::Alias(ty) => ty.docs(),
106 Self::Enum(ty) => ty.docs(),
107 }
108 }
109
110 pub fn ty(&self) -> TypeExpr {
112 match self {
113 Self::Alias(ty) => ty.ty.clone(),
114 Self::Enum(ty) => TypeExpr::Primitive(Span::new(ty.span, ty.ty.clone())),
115 }
116 }
117}
118
119impl Spanned for TypeDecl {
120 fn span(&self) -> SourceSpan {
121 match self {
122 Self::Alias(spanned) => spanned.span,
123 Self::Enum(spanned) => spanned.span,
124 }
125 }
126}
127
128impl From<TypeAlias> for TypeDecl {
129 fn from(value: TypeAlias) -> Self {
130 Self::Alias(value)
131 }
132}
133
134impl From<EnumType> for TypeDecl {
135 fn from(value: EnumType) -> Self {
136 Self::Enum(value)
137 }
138}
139
140impl crate::prettier::PrettyPrint for TypeDecl {
141 fn render(&self) -> crate::prettier::Document {
142 match self {
143 Self::Alias(ty) => ty.render(),
144 Self::Enum(ty) => ty.render(),
145 }
146 }
147}
148
149#[derive(Debug, Clone)]
154pub struct FunctionType {
155 pub span: SourceSpan,
156 pub cc: types::CallConv,
157 pub args: Vec<TypeExpr>,
158 pub arg_names: Vec<Option<Ident>>,
160 pub results: Vec<TypeExpr>,
161}
162
163impl Eq for FunctionType {}
164
165impl PartialEq for FunctionType {
166 fn eq(&self, other: &Self) -> bool {
167 self.cc == other.cc && self.args == other.args && self.results == other.results
168 }
169}
170
171impl core::hash::Hash for FunctionType {
172 fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
173 self.cc.hash(state);
174 self.args.hash(state);
175 self.results.hash(state);
176 }
177}
178
179impl Spanned for FunctionType {
180 fn span(&self) -> SourceSpan {
181 self.span
182 }
183}
184
185impl FunctionType {
186 pub fn new(cc: types::CallConv, args: Vec<TypeExpr>, results: Vec<TypeExpr>) -> Self {
187 Self {
188 span: SourceSpan::UNKNOWN,
189 cc,
190 args,
191 arg_names: Vec::new(),
192 results,
193 }
194 }
195
196 #[inline]
198 pub fn with_span(mut self, span: SourceSpan) -> Self {
199 self.span = span;
200 self
201 }
202
203 #[inline]
205 pub fn with_arg_names(mut self, arg_names: Vec<Option<Ident>>) -> Self {
206 debug_assert_eq!(arg_names.len(), self.args.len());
207 self.arg_names = arg_names;
208 self
209 }
210}
211
212impl crate::prettier::PrettyPrint for FunctionType {
213 fn render(&self) -> crate::prettier::Document {
214 use crate::prettier::*;
215
216 let render_arg = |(index, ty): (usize, &TypeExpr)| {
217 if matches!(ty, TypeExpr::Primitive(prim) if matches!(prim.inner(), Type::Variadic)) {
218 return ty.render();
219 }
220 let name = match self.arg_names.get(index) {
221 Some(Some(name)) => display(name),
222 _ => text(format!("arg{index}")),
223 };
224 name + const_text(": ") + ty.render()
225 };
226 let singleline_args = self
227 .args
228 .iter()
229 .enumerate()
230 .map(render_arg)
231 .reduce(|acc, arg| acc + const_text(", ") + arg)
232 .unwrap_or(Document::Empty);
233 let multiline_args = indent(
234 4,
235 nl() + self
236 .args
237 .iter()
238 .enumerate()
239 .map(render_arg)
240 .reduce(|acc, arg| acc + const_text(",") + nl() + arg)
241 .unwrap_or(Document::Empty),
242 ) + nl();
243 let args = singleline_args | multiline_args;
244 let args = const_text("(") + args + const_text(")");
245
246 match self.results.len() {
247 0 => args,
248 1 => args + const_text(" -> ") + self.results[0].render(),
249 _ => {
250 let results = self
251 .results
252 .iter()
253 .map(PrettyPrint::render)
254 .reduce(|acc, r| acc + const_text(", ") + r)
255 .unwrap_or(Document::Empty);
256 args + const_text(" -> ") + const_text("(") + results + const_text(")")
257 },
258 }
259 }
260}
261
262#[derive(Debug, Clone, Eq, PartialEq, Hash)]
267pub enum TypeExpr {
268 Primitive(Span<Type>),
270 Ptr(PointerType),
272 Array(ArrayType),
274 Struct(StructType),
276 Ref(Span<Arc<Path>>),
278}
279
280impl TypeExpr {
281 pub fn set_name(&mut self, name: Ident) {
286 match self {
287 Self::Struct(struct_ty) => {
288 struct_ty.name = Some(name);
289 },
290 Self::Primitive(_) | Self::Ptr(_) | Self::Array(_) | Self::Ref(_) => (),
291 }
292 }
293
294 pub fn references(&self) -> Vec<Span<Arc<Path>>> {
296 use alloc::collections::BTreeSet;
297
298 let mut worklist = smallvec::SmallVec::<[_; 4]>::from_slice(&[self]);
299 let mut references = BTreeSet::new();
300
301 while let Some(ty) = worklist.pop() {
302 match ty {
303 Self::Primitive(_) => {},
304 Self::Ptr(ty) => {
305 worklist.push(&ty.pointee);
306 },
307 Self::Array(ty) => {
308 worklist.push(&ty.elem);
309 },
310 Self::Struct(ty) => {
311 for field in ty.fields.iter() {
312 worklist.push(&field.ty);
313 }
314 },
315 Self::Ref(ty) => {
316 references.insert(ty.clone());
317 },
318 }
319 }
320
321 references.into_iter().collect()
322 }
323
324 pub fn resolve_template<E, R>(&self, resolver: &mut R) -> Result<Option<TypeTemplate>, E>
328 where
329 R: ?Sized + TypeResolver<E>,
330 {
331 self.resolve_template_with_depth(resolver, 0)
332 }
333
334 fn resolve_template_with_depth<E, R>(
335 &self,
336 resolver: &mut R,
337 depth: usize,
338 ) -> Result<Option<TypeTemplate>, E>
339 where
340 R: ?Sized + TypeResolver<E>,
341 {
342 if depth > MAX_TYPE_EXPR_NESTING {
343 let source_manager = resolver.source_manager();
344 return Err(resolver.resolve_local_failed(
345 SymbolResolutionError::type_expression_depth_exceeded(
346 self.span(),
347 MAX_TYPE_EXPR_NESTING,
348 source_manager.as_ref(),
349 ),
350 ));
351 }
352
353 match self {
354 TypeExpr::Ref(path) => {
355 let mut current_path = path.clone();
356 loop {
357 match resolver.resolve_type_ref(current_path.as_deref())? {
358 SymbolResolution::Local(item) => {
359 return resolver.get_local_type(current_path.span(), item.into_inner());
360 },
361 SymbolResolution::External(path) => {
362 if path == current_path {
364 break Ok(None);
365 }
366 current_path = path;
367 },
368 SymbolResolution::Exact { gid, .. } => {
369 return resolver.get_type(current_path.span(), gid);
370 },
371 SymbolResolution::Module { path: module_path, .. } => {
372 break Err(resolver.resolve_local_failed(
373 SymbolResolutionError::invalid_symbol_type(
374 path.span(),
375 "type",
376 module_path.span(),
377 &resolver.source_manager(),
378 ),
379 ));
380 },
381 SymbolResolution::MastRoot(item) => {
382 break Err(resolver.resolve_local_failed(
383 SymbolResolutionError::invalid_symbol_type(
384 path.span(),
385 "type",
386 item.span(),
387 &resolver.source_manager(),
388 ),
389 ));
390 },
391 }
392 }
393 },
394 TypeExpr::Primitive(t) => Ok(Some(TypeTemplate::Type(t.inner().clone()))),
395 TypeExpr::Array(t) => Ok(t
396 .elem
397 .resolve_template_with_depth(resolver, depth + 1)?
398 .map(|elem| TypeTemplate::array(elem, t.arity))),
399 TypeExpr::Ptr(ty) => Ok(ty
400 .pointee
401 .resolve_template_with_depth(resolver, depth + 1)?
402 .map(|pointee| TypeTemplate::ptr_in(ty.address_space(), pointee))),
403 TypeExpr::Struct(t) => {
404 let mut fields = Vec::with_capacity(t.fields.len());
405 for field in t.fields.iter() {
406 let field_ty = field.ty.resolve_template_with_depth(resolver, depth + 1)?;
407 if let Some(field_ty) = field_ty {
408 fields.push(types::FieldTemplate {
409 name: Some(field.name.clone().into_inner()),
410 ty: field_ty,
411 });
412 } else {
413 return Ok(None);
414 }
415 }
416 Ok(Some(TypeTemplate::Struct(Box::new(types::StructTemplate {
417 name: t.name.clone().map(Ident::into_inner),
418 repr: t.repr.into_inner(),
419 fields,
420 }))))
421 },
422 }
423 }
424}
425
426impl From<Type> for TypeExpr {
427 fn from(ty: Type) -> Self {
428 let mut expanding = Vec::new();
429 type_expr_from(ty, &mut expanding)
430 }
431}
432
433fn type_expr_from(ty: Type, expanding: &mut Vec<types::RecTypeRef>) -> TypeExpr {
440 match ty {
441 Type::Array(t) => TypeExpr::Array(ArrayType::new(
442 type_expr_from(t.element_type().clone(), expanding),
443 t.len(),
444 )),
445 Type::Struct(t) => {
446 let name = t.name().and_then(|name| Ident::new(name.as_ref()).ok());
447
448 if let Some(rec) = t.as_recursive() {
451 if expanding.contains(rec) {
452 let name = name.unwrap_or_else(|| {
453 panic!(
454 "unrepresentable type value: a recursive struct without a name cannot \
455 be referred to as a type expression"
456 )
457 });
458 return TypeExpr::Ref(Span::unknown(
459 Path::from_ident(&name).into_owned().into(),
460 ));
461 }
462 expanding.push(rec.clone());
463 }
464
465 let is_recursive = t.is_recursive();
466 let body = t.get();
467 let fields = body
468 .fields()
469 .iter()
470 .enumerate()
471 .map(|(i, ft)| {
472 let name = ft
473 .name
474 .as_deref()
475 .map(Ident::new)
476 .and_then(Result::ok)
477 .unwrap_or_else(|| Ident::new(format!("field{i}")).unwrap());
478 StructField {
479 span: SourceSpan::UNKNOWN,
480 name,
481 ty: type_expr_from(ft.ty.clone(), expanding),
482 }
483 })
484 .collect::<Vec<_>>();
485 let converted = TypeExpr::Struct(
486 StructType::new(name, fields)
487 .with_repr(Span::unknown(body.repr()))
488 .with_span(SourceSpan::UNKNOWN),
489 );
490
491 if is_recursive {
492 expanding.pop();
493 }
494 converted
495 },
496 Type::Ptr(t) => TypeExpr::Ptr(
497 PointerType::new(type_expr_from(t.pointee().clone(), expanding))
498 .with_address_space(t.addrspace()),
499 ),
500 Type::Function(_) => {
501 TypeExpr::Ptr(PointerType::new(TypeExpr::Primitive(Span::unknown(Type::Felt))))
502 },
503 Type::List(t) => TypeExpr::Ptr(
504 PointerType::new(type_expr_from((*t).clone(), expanding))
505 .with_address_space(AddressSpace::Byte),
506 ),
507 Type::Unknown | Type::Never | Type::F64 => {
508 panic!("unrepresentable type value: {ty}")
509 },
510 ty => TypeExpr::Primitive(Span::unknown(ty)),
511 }
512}
513
514impl Spanned for TypeExpr {
515 fn span(&self) -> SourceSpan {
516 match self {
517 Self::Primitive(spanned) => spanned.span(),
518 Self::Ptr(spanned) => spanned.span(),
519 Self::Array(spanned) => spanned.span(),
520 Self::Struct(spanned) => spanned.span(),
521 Self::Ref(spanned) => spanned.span(),
522 }
523 }
524}
525
526impl crate::prettier::PrettyPrint for TypeExpr {
527 fn render(&self) -> crate::prettier::Document {
528 use crate::prettier::*;
529
530 match self {
531 Self::Primitive(ty) => display(ty),
532 Self::Ptr(ty) => ty.render(),
533 Self::Array(ty) => ty.render(),
534 Self::Struct(ty) => ty.render(),
535 Self::Ref(ty) => display(ty),
536 }
537 }
538}
539
540#[derive(Debug, Clone)]
544pub struct PointerType {
545 pub span: SourceSpan,
546 pub pointee: Box<TypeExpr>,
547 addrspace: Option<AddressSpace>,
548}
549
550impl From<types::PointerType> for PointerType {
551 fn from(ty: types::PointerType) -> Self {
552 let types::PointerType { addrspace, pointee } = ty;
553 let pointee = Box::new(TypeExpr::from(pointee));
554 Self {
555 span: SourceSpan::UNKNOWN,
556 pointee,
557 addrspace: Some(addrspace),
558 }
559 }
560}
561
562impl Eq for PointerType {}
563
564impl PartialEq for PointerType {
565 fn eq(&self, other: &Self) -> bool {
566 self.address_space() == other.address_space() && self.pointee == other.pointee
567 }
568}
569
570impl core::hash::Hash for PointerType {
571 fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
572 self.pointee.hash(state);
573 self.address_space().hash(state);
574 }
575}
576
577impl Spanned for PointerType {
578 fn span(&self) -> SourceSpan {
579 self.span
580 }
581}
582
583impl PointerType {
584 pub fn new(pointee: TypeExpr) -> Self {
585 Self {
586 span: SourceSpan::UNKNOWN,
587 pointee: Box::new(pointee),
588 addrspace: None,
589 }
590 }
591
592 #[inline]
594 pub fn with_span(mut self, span: SourceSpan) -> Self {
595 self.span = span;
596 self
597 }
598
599 #[inline]
601 pub fn with_address_space(mut self, addrspace: AddressSpace) -> Self {
602 self.addrspace = Some(addrspace);
603 self
604 }
605
606 #[inline]
608 pub fn address_space(&self) -> AddressSpace {
609 self.addrspace.unwrap_or(AddressSpace::Element)
610 }
611}
612
613impl crate::prettier::PrettyPrint for PointerType {
614 fn render(&self) -> crate::prettier::Document {
615 use crate::prettier::*;
616
617 let doc = const_text("ptr<") + self.pointee.render();
618 if let Some(addrspace) = self.addrspace.as_ref() {
619 doc + const_text(", ") + text(format!("addrspace({addrspace})")) + const_text(">")
620 } else {
621 doc + const_text(">")
622 }
623 }
624}
625
626#[derive(Debug, Clone)]
630pub struct ArrayType {
631 pub span: SourceSpan,
632 pub elem: Box<TypeExpr>,
633 pub arity: usize,
634}
635
636impl Eq for ArrayType {}
637
638impl PartialEq for ArrayType {
639 fn eq(&self, other: &Self) -> bool {
640 self.arity == other.arity && self.elem == other.elem
641 }
642}
643
644impl core::hash::Hash for ArrayType {
645 fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
646 self.elem.hash(state);
647 self.arity.hash(state);
648 }
649}
650
651impl Spanned for ArrayType {
652 fn span(&self) -> SourceSpan {
653 self.span
654 }
655}
656
657impl ArrayType {
658 pub fn new(elem: TypeExpr, arity: usize) -> Self {
659 Self {
660 span: SourceSpan::UNKNOWN,
661 elem: Box::new(elem),
662 arity,
663 }
664 }
665
666 #[inline]
668 pub fn with_span(mut self, span: SourceSpan) -> Self {
669 self.span = span;
670 self
671 }
672}
673
674impl crate::prettier::PrettyPrint for ArrayType {
675 fn render(&self) -> crate::prettier::Document {
676 use crate::prettier::*;
677
678 const_text("[")
679 + self.elem.render()
680 + const_text("; ")
681 + display(self.arity)
682 + const_text("]")
683 }
684}
685
686#[derive(Debug, Clone)]
690pub struct StructType {
691 pub span: SourceSpan,
692 pub name: Option<Ident>,
693 pub repr: Span<TypeRepr>,
694 pub fields: Vec<StructField>,
695}
696
697impl Eq for StructType {}
698
699impl PartialEq for StructType {
700 fn eq(&self, other: &Self) -> bool {
701 self.name == other.name && self.repr == other.repr && self.fields == other.fields
702 }
703}
704
705impl core::hash::Hash for StructType {
706 fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
707 self.name.hash(state);
708 self.repr.hash(state);
709 self.fields.hash(state);
710 }
711}
712
713impl Spanned for StructType {
714 fn span(&self) -> SourceSpan {
715 self.span
716 }
717}
718
719impl StructType {
720 pub fn new(name: Option<Ident>, fields: impl IntoIterator<Item = StructField>) -> Self {
721 Self {
722 span: SourceSpan::UNKNOWN,
723 name,
724 repr: Span::unknown(TypeRepr::Default),
725 fields: fields.into_iter().collect(),
726 }
727 }
728
729 #[inline]
731 pub fn with_repr(mut self, repr: Span<TypeRepr>) -> Self {
732 self.repr = repr;
733 self
734 }
735
736 #[inline]
738 pub fn with_span(mut self, span: SourceSpan) -> Self {
739 self.span = span;
740 self
741 }
742}
743
744impl crate::prettier::PrettyPrint for StructType {
745 fn render(&self) -> crate::prettier::Document {
746 use crate::prettier::*;
747
748 let repr = match &*self.repr {
749 TypeRepr::Default => Document::Empty,
750 repr @ (TypeRepr::Align(_) | TypeRepr::Packed(_) | TypeRepr::Transparent) => {
751 text(format!(" @{repr}"))
752 },
753 };
754
755 let singleline_body = self
756 .fields
757 .iter()
758 .map(PrettyPrint::render)
759 .reduce(|acc, field| acc + const_text(", ") + field)
760 .unwrap_or(Document::Empty);
761 let multiline_body = indent(
762 4,
763 nl() + self
764 .fields
765 .iter()
766 .map(PrettyPrint::render)
767 .reduce(|acc, field| acc + const_text(",") + nl() + field)
768 .unwrap_or(Document::Empty),
769 ) + nl();
770 let body = singleline_body | multiline_body;
771
772 const_text("struct") + repr + const_text(" { ") + body + const_text(" }")
773 }
774}
775
776#[derive(Debug, Clone)]
780pub struct StructField {
781 pub span: SourceSpan,
782 pub name: Ident,
783 pub ty: TypeExpr,
784}
785
786impl Eq for StructField {}
787
788impl PartialEq for StructField {
789 fn eq(&self, other: &Self) -> bool {
790 self.name == other.name && self.ty == other.ty
791 }
792}
793
794impl core::hash::Hash for StructField {
795 fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
796 self.name.hash(state);
797 self.ty.hash(state);
798 }
799}
800
801impl Spanned for StructField {
802 fn span(&self) -> SourceSpan {
803 self.span
804 }
805}
806
807impl crate::prettier::PrettyPrint for StructField {
808 fn render(&self) -> crate::prettier::Document {
809 use crate::prettier::*;
810
811 display(&self.name) + const_text(": ") + self.ty.render()
812 }
813}
814
815#[derive(Debug, Clone)]
824pub struct TypeAlias {
825 span: SourceSpan,
826 docs: Option<DocString>,
828 pub visibility: Visibility,
830 pub name: Ident,
832 pub ty: TypeExpr,
834}
835
836impl TypeAlias {
837 pub fn new(visibility: Visibility, name: Ident, ty: TypeExpr) -> Self {
839 Self {
840 span: name.span(),
841 docs: None,
842 visibility,
843 name,
844 ty,
845 }
846 }
847
848 pub fn with_docs(mut self, docs: Option<Span<String>>) -> Self {
850 self.docs = docs.map(DocString::new);
851 self
852 }
853
854 #[inline]
856 pub fn with_span(mut self, span: SourceSpan) -> Self {
857 self.span = span;
858 self
859 }
860
861 #[inline]
863 pub fn set_span(&mut self, span: SourceSpan) {
864 self.span = span;
865 }
866
867 pub fn docs(&self) -> Option<Span<&str>> {
869 self.docs.as_ref().map(|docstring| docstring.as_spanned_str())
870 }
871
872 pub fn name(&self) -> &Ident {
874 &self.name
875 }
876
877 #[inline]
879 pub const fn visibility(&self) -> Visibility {
880 self.visibility
881 }
882}
883
884impl Eq for TypeAlias {}
885
886impl PartialEq for TypeAlias {
887 fn eq(&self, other: &Self) -> bool {
888 self.visibility == other.visibility
889 && self.name == other.name
890 && self.docs == other.docs
891 && self.ty == other.ty
892 }
893}
894
895impl core::hash::Hash for TypeAlias {
896 fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
897 let Self { span: _, docs, visibility, name, ty } = self;
898 docs.hash(state);
899 visibility.hash(state);
900 name.hash(state);
901 ty.hash(state);
902 }
903}
904
905impl Spanned for TypeAlias {
906 fn span(&self) -> SourceSpan {
907 self.span
908 }
909}
910
911impl crate::prettier::PrettyPrint for TypeAlias {
912 fn render(&self) -> crate::prettier::Document {
913 use crate::prettier::*;
914
915 let mut doc = self.docs.as_ref().map(PrettyPrint::render).unwrap_or(Document::Empty);
916
917 if self.visibility.is_public() {
918 doc += display(self.visibility) + const_text(" ");
919 }
920
921 doc + const_text("type")
922 + const_text(" ")
923 + display(&self.name)
924 + const_text(" = ")
925 + self.ty.render()
926 }
927}
928
929#[derive(Debug, Clone)]
943pub struct EnumType {
944 span: SourceSpan,
945 docs: Option<DocString>,
947 visibility: Visibility,
949 name: Ident,
951 ty: Type,
955 variants: Vec<Variant>,
957}
958
959impl EnumType {
960 pub fn new(
965 visibility: Visibility,
966 name: Ident,
967 ty: Type,
968 variants: impl IntoIterator<Item = Variant>,
969 ) -> Self {
970 assert!(ty.is_integer(), "only integer types are allowed in enum type definitions");
971 Self {
972 span: name.span(),
973 docs: None,
974 visibility,
975 name,
976 ty,
977 variants: Vec::from_iter(variants),
978 }
979 }
980
981 pub fn with_docs(mut self, docs: Option<Span<String>>) -> Self {
983 self.docs = docs.map(DocString::new);
984 self
985 }
986
987 pub fn with_span(mut self, span: SourceSpan) -> Self {
989 self.span = span;
990 self
991 }
992
993 pub fn is_c_like(&self) -> bool {
995 !self.variants.is_empty() && self.variants.iter().all(|v| v.value_ty.is_none())
996 }
997
998 pub fn set_span(&mut self, span: SourceSpan) {
1000 self.span = span;
1001 }
1002
1003 pub fn name(&self) -> &Ident {
1005 &self.name
1006 }
1007
1008 pub const fn visibility(&self) -> Visibility {
1010 self.visibility
1011 }
1012
1013 pub fn docs(&self) -> Option<Span<&str>> {
1015 self.docs.as_ref().map(|docstring| docstring.as_spanned_str())
1016 }
1017
1018 pub fn ty(&self) -> &Type {
1020 &self.ty
1021 }
1022
1023 pub fn variants(&self) -> &[Variant] {
1025 &self.variants
1026 }
1027
1028 pub fn variants_mut(&mut self) -> &mut Vec<Variant> {
1030 &mut self.variants
1031 }
1032
1033 pub fn into_parts(self) -> (TypeAlias, Vec<Variant>) {
1035 let Self {
1036 span,
1037 docs,
1038 visibility,
1039 name,
1040 ty,
1041 variants,
1042 } = self;
1043 let alias = TypeAlias {
1044 span,
1045 docs,
1046 visibility,
1047 name,
1048 ty: TypeExpr::Primitive(Span::new(span, ty)),
1049 };
1050 (alias, variants)
1051 }
1052}
1053
1054impl Spanned for EnumType {
1055 fn span(&self) -> SourceSpan {
1056 self.span
1057 }
1058}
1059
1060impl Eq for EnumType {}
1061
1062impl PartialEq for EnumType {
1063 fn eq(&self, other: &Self) -> bool {
1064 self.visibility == other.visibility
1065 && self.name == other.name
1066 && self.docs == other.docs
1067 && self.ty == other.ty
1068 && self.variants == other.variants
1069 }
1070}
1071
1072impl core::hash::Hash for EnumType {
1073 fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
1074 let Self {
1075 span: _,
1076 docs,
1077 visibility,
1078 name,
1079 ty,
1080 variants,
1081 } = self;
1082 docs.hash(state);
1083 visibility.hash(state);
1084 name.hash(state);
1085 ty.hash(state);
1086 variants.hash(state);
1087 }
1088}
1089
1090impl crate::prettier::PrettyPrint for EnumType {
1091 fn render(&self) -> crate::prettier::Document {
1092 use crate::prettier::*;
1093
1094 let mut doc = self.docs.as_ref().map(PrettyPrint::render).unwrap_or(Document::Empty);
1095
1096 let variants = self
1097 .variants
1098 .iter()
1099 .map(PrettyPrint::render)
1100 .reduce(|acc, v| acc + const_text(",") + nl() + v)
1101 .unwrap_or(Document::Empty);
1102
1103 if self.visibility.is_public() {
1104 doc += display(self.visibility) + const_text(" ");
1105 }
1106
1107 doc + const_text("enum")
1108 + const_text(" ")
1109 + display(&self.name)
1110 + const_text(" : ")
1111 + self.ty.render()
1112 + const_text(" {")
1113 + nl()
1114 + variants
1115 + const_text("}")
1116 }
1117}
1118
1119#[derive(Debug, Clone)]
1126pub struct Variant {
1127 pub span: SourceSpan,
1128 pub docs: Option<DocString>,
1130 pub name: Ident,
1132 pub value_ty: Option<TypeExpr>,
1137 pub discriminant: ConstantExpr,
1139}
1140
1141impl Variant {
1142 pub fn new(name: Ident, discriminant: ConstantExpr, payload: Option<TypeExpr>) -> Self {
1144 Self {
1145 span: name.span(),
1146 docs: None,
1147 name,
1148 value_ty: payload,
1149 discriminant,
1150 }
1151 }
1152
1153 pub fn with_span(mut self, span: SourceSpan) -> Self {
1155 self.span = span;
1156 self
1157 }
1158
1159 pub fn with_docs(mut self, docs: Option<Span<String>>) -> Self {
1161 self.docs = docs.map(DocString::new);
1162 self
1163 }
1164
1165 pub fn assert_instance_of(&self, ty: &Type) -> Result<(), crate::SemanticAnalysisError> {
1173 use crate::{FIELD_MODULUS, SemanticAnalysisError};
1174
1175 let value = match &self.discriminant {
1176 ConstantExpr::Int(value) => value.as_int(),
1177 _ => {
1178 return Err(SemanticAnalysisError::InvalidEnumDiscriminant {
1179 span: self.discriminant.span(),
1180 repr: ty.clone(),
1181 });
1182 },
1183 };
1184
1185 match ty {
1186 Type::Felt if value >= FIELD_MODULUS => {
1187 Err(SemanticAnalysisError::InvalidEnumDiscriminant {
1188 span: self.discriminant.span(),
1189 repr: ty.clone(),
1190 })
1191 },
1192 Type::Felt => Ok(()),
1195 Type::I1 if value > 1 => Err(SemanticAnalysisError::InvalidEnumDiscriminant {
1196 span: self.discriminant.span(),
1197 repr: ty.clone(),
1198 }),
1199 Type::I1 => Ok(()),
1200 Type::I8 | Type::U8 if value > u8::MAX as u64 => {
1201 Err(SemanticAnalysisError::InvalidEnumDiscriminant {
1202 span: self.discriminant.span(),
1203 repr: ty.clone(),
1204 })
1205 },
1206 Type::I8 | Type::U8 => Ok(()),
1207 Type::I16 | Type::U16 if value > u16::MAX as u64 => {
1208 Err(SemanticAnalysisError::InvalidEnumDiscriminant {
1209 span: self.discriminant.span(),
1210 repr: ty.clone(),
1211 })
1212 },
1213 Type::I16 | Type::U16 => Ok(()),
1214 Type::I32 | Type::U32 if value > u32::MAX as u64 => {
1215 Err(SemanticAnalysisError::InvalidEnumDiscriminant {
1216 span: self.discriminant.span(),
1217 repr: ty.clone(),
1218 })
1219 },
1220 Type::I32 | Type::U32 => Ok(()),
1221 Type::I64 | Type::U64 if value >= FIELD_MODULUS => {
1222 Err(SemanticAnalysisError::InvalidEnumDiscriminant {
1223 span: self.discriminant.span(),
1224 repr: ty.clone(),
1225 })
1226 },
1227 _ => Err(SemanticAnalysisError::InvalidEnumRepr { span: self.span }),
1228 }
1229 }
1230}
1231
1232impl Spanned for Variant {
1233 fn span(&self) -> SourceSpan {
1234 self.span
1235 }
1236}
1237
1238impl Eq for Variant {}
1239
1240impl PartialEq for Variant {
1241 fn eq(&self, other: &Self) -> bool {
1242 self.name == other.name
1243 && self.value_ty == other.value_ty
1244 && self.discriminant == other.discriminant
1245 && self.docs == other.docs
1246 }
1247}
1248
1249impl core::hash::Hash for Variant {
1250 fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
1251 let Self {
1252 span: _,
1253 docs,
1254 name,
1255 value_ty,
1256 discriminant,
1257 } = self;
1258 docs.hash(state);
1259 name.hash(state);
1260 value_ty.hash(state);
1261 discriminant.hash(state);
1262 }
1263}
1264
1265impl crate::prettier::PrettyPrint for Variant {
1266 fn render(&self) -> crate::prettier::Document {
1267 use crate::prettier::*;
1268
1269 let doc = self.docs.as_ref().map(PrettyPrint::render).unwrap_or(Document::Empty);
1270
1271 let name = display(&self.name);
1272 let name_and_payload = if let Some(value_ty) = self.value_ty.as_ref() {
1273 name + const_text("(") + value_ty.render() + const_text(")")
1274 } else {
1275 name
1276 };
1277 doc + name_and_payload + const_text(" = ") + self.discriminant.render()
1278 }
1279}
1280
1281#[cfg(test)]
1282mod tests {
1283 use alloc::{string::ToString, sync::Arc};
1284 use core::str::FromStr;
1285
1286 use miden_debug_types::{DefaultSourceManager, SourceFile, SourceId, SourceLanguage, Uri};
1287
1288 use super::*;
1289 use crate::{ast::Form, prettier::PrettyPrint};
1290
1291 struct DummyResolver {
1292 source_manager: Arc<dyn SourceManager>,
1293 }
1294
1295 impl DummyResolver {
1296 fn new() -> Self {
1297 Self {
1298 source_manager: Arc::new(DefaultSourceManager::default()),
1299 }
1300 }
1301 }
1302
1303 impl TypeResolver<SymbolResolutionError> for DummyResolver {
1304 fn source_manager(&self) -> Arc<dyn SourceManager> {
1305 self.source_manager.clone()
1306 }
1307
1308 fn resolve_local_failed(&self, err: SymbolResolutionError) -> SymbolResolutionError {
1309 err
1310 }
1311
1312 fn get_type(
1313 &mut self,
1314 context: SourceSpan,
1315 _gid: GlobalItemIndex,
1316 ) -> Result<Option<TypeTemplate>, SymbolResolutionError> {
1317 Err(SymbolResolutionError::undefined(context, self.source_manager.as_ref()))
1318 }
1319
1320 fn get_local_type(
1321 &mut self,
1322 _context: SourceSpan,
1323 _id: ItemIndex,
1324 ) -> Result<Option<TypeTemplate>, SymbolResolutionError> {
1325 Ok(None)
1326 }
1327
1328 fn resolve_type_ref(
1329 &mut self,
1330 ty: Span<&Path>,
1331 ) -> Result<SymbolResolution, SymbolResolutionError> {
1332 Err(SymbolResolutionError::undefined(ty.span(), self.source_manager.as_ref()))
1333 }
1334
1335 fn finalize(
1336 &mut self,
1337 context: SourceSpan,
1338 template: TypeTemplate,
1339 ) -> Result<Type, SymbolResolutionError> {
1340 midenc_hir_type::close_template(&template, |_| None).map_err(|_| {
1342 SymbolResolutionError::undefined(context, self.source_manager.as_ref())
1343 })
1344 }
1345 }
1346
1347 fn nested_type_expr(depth: usize) -> TypeExpr {
1348 let mut expr = TypeExpr::Primitive(Span::unknown(Type::Felt));
1349 for i in 0..depth {
1350 expr = match i % 3 {
1351 0 => TypeExpr::Ptr(PointerType::new(expr)),
1352 1 => TypeExpr::Array(ArrayType::new(expr, 1)),
1353 _ => {
1354 let field = StructField {
1355 span: SourceSpan::UNKNOWN,
1356 name: Ident::from_str("field").expect("valid ident"),
1357 ty: expr,
1358 };
1359 TypeExpr::Struct(StructType::new(None, [field]))
1360 },
1361 };
1362 }
1363 expr
1364 }
1365
1366 fn test_source_file(source: &str) -> Arc<SourceFile> {
1367 Arc::new(SourceFile::new(
1368 SourceId::default(),
1369 SourceLanguage::Masm,
1370 Uri::new("memory:///type-expr-test.masm"),
1371 source.to_string().into_boxed_str(),
1372 ))
1373 }
1374
1375 fn parse_type_alias_expr(source: &str) -> TypeExpr {
1376 let mut forms =
1377 crate::parser::parse_forms(test_source_file(source)).expect("type alias should parse");
1378 assert_eq!(forms.len(), 1, "expected exactly one parsed form");
1379 match forms.pop().expect("expected parsed form") {
1380 Form::Type(alias) => alias.ty,
1381 form => panic!("expected type alias form, got {form:?}"),
1382 }
1383 }
1384
1385 fn repr_round_trip_struct(repr: TypeRepr) -> TypeExpr {
1386 TypeExpr::Struct(
1387 StructType::new(
1388 None,
1389 [
1390 StructField {
1391 span: SourceSpan::UNKNOWN,
1392 name: Ident::from_str("prefix").expect("valid ident"),
1393 ty: TypeExpr::Primitive(Span::unknown(Type::Felt)),
1394 },
1395 StructField {
1396 span: SourceSpan::UNKNOWN,
1397 name: Ident::from_str("suffix").expect("valid ident"),
1398 ty: TypeExpr::Primitive(Span::unknown(Type::U32)),
1399 },
1400 ],
1401 )
1402 .with_repr(Span::unknown(repr)),
1403 )
1404 }
1405
1406 #[test]
1407 fn type_expr_depth_boundary() {
1408 let mut resolver = DummyResolver::new();
1409
1410 let ok_expr = nested_type_expr(MAX_TYPE_EXPR_NESTING);
1411 assert!(ok_expr.resolve_template(&mut resolver).is_ok());
1412
1413 let err_expr = nested_type_expr(MAX_TYPE_EXPR_NESTING + 1);
1414 let err = err_expr
1415 .resolve_template(&mut resolver)
1416 .expect_err("expected depth-exceeded error");
1417 assert!(
1418 matches!(err, SymbolResolutionError::TypeExpressionDepthExceeded { max_depth, .. }
1419 if max_depth == MAX_TYPE_EXPR_NESTING)
1420 );
1421 }
1422
1423 #[test]
1424 fn struct_type_expr_render_round_trips_non_default_reprs() {
1425 for repr in [
1426 TypeRepr::align(16),
1427 TypeRepr::packed(1),
1428 TypeRepr::packed(2),
1429 TypeRepr::Transparent,
1430 ] {
1431 let rendered = repr_round_trip_struct(repr).to_pretty_string();
1432 assert!(
1433 rendered.starts_with("struct @"),
1434 "non-default struct repr should render after `struct`: {rendered}"
1435 );
1436
1437 let parsed = parse_type_alias_expr(&format!("type RoundTrip = {rendered}\n"));
1438 let TypeExpr::Struct(parsed) = parsed else {
1439 panic!("expected rendered type to parse back as a struct");
1440 };
1441 assert_eq!(*parsed.repr, repr);
1442 assert_eq!(parsed.fields[0].name.as_str(), "prefix");
1443 assert_eq!(parsed.fields[1].name.as_str(), "suffix");
1444 }
1445 }
1446
1447 #[test]
1448 fn type_expr_from_type_preserves_wide_integer_primitives() {
1449 for ty in [Type::I64, Type::U64, Type::I128, Type::U128] {
1450 let expr = TypeExpr::from(ty.clone());
1451 let TypeExpr::Primitive(actual) = expr else {
1452 panic!("expected primitive type expression for {ty}, got {expr:?}");
1453 };
1454 assert_eq!(actual.into_inner(), ty);
1455 }
1456 }
1457
1458 #[test]
1459 fn type_expr_from_type_preserves_struct_metadata() {
1460 let ty = Type::from(Arc::new(types::StructType::from_parts(
1461 Some(Arc::from("miden:base/core-types@1.0.0/account-id")),
1462 TypeRepr::align(16),
1463 [
1464 (Arc::<str>::from("prefix"), Type::Felt),
1465 (Arc::<str>::from("suffix"), Type::Felt),
1466 ],
1467 )));
1468
1469 let TypeExpr::Struct(actual) = TypeExpr::from(ty) else {
1470 panic!("expected struct type expression");
1471 };
1472 assert_eq!(
1473 actual.name.as_ref().map(Ident::as_str),
1474 Some("miden:base/core-types@1.0.0/account-id"),
1475 );
1476 assert_eq!(*actual.repr, TypeRepr::align(16));
1477 assert_eq!(actual.fields[0].name.as_str(), "prefix");
1478 assert_eq!(actual.fields[1].name.as_str(), "suffix");
1479 }
1480
1481 #[test]
1482 fn type_expr_conversion_of_a_recursive_struct_terminates() {
1483 use midenc_hir_type::{RecursiveTypeBuilder, StructTemplate, TypeRepr, TypeTemplate};
1484
1485 let mut builder = RecursiveTypeBuilder::new();
1486 builder.define_struct(
1487 "Node",
1488 StructTemplate::named(
1489 "Node",
1490 TypeRepr::Default,
1491 [("next", TypeTemplate::ptr(TypeTemplate::rec("Node")))],
1492 ),
1493 );
1494 let node = builder.build().unwrap().remove("Node").unwrap();
1495
1496 let TypeExpr::Struct(converted) = TypeExpr::from(node) else {
1499 panic!("expected a struct type expression");
1500 };
1501 let TypeExpr::Ptr(pointer) = &converted.fields[0].ty else {
1502 panic!("expected a pointer");
1503 };
1504 let TypeExpr::Ref(target) = pointer.pointee.as_ref() else {
1505 panic!("expected the pointee to be a reference, got {:?}", pointer.pointee);
1506 };
1507 assert_eq!(target.inner().to_string(), "Node");
1508 }
1509
1510 #[test]
1511 fn parsed_struct_type_preserves_field_names_through_resolution() {
1512 let expr = parse_type_alias_expr(
1513 "type AccountId = struct @align(16) { prefix: felt, suffix: felt }\n",
1514 );
1515
1516 let mut resolver = DummyResolver::new();
1517 let resolved = TypeResolver::resolve(&mut resolver, &expr)
1518 .expect("struct type should resolve")
1519 .expect("struct type should be concrete");
1520 let Type::Struct(resolved_struct) = &resolved else {
1521 panic!("expected resolved struct type, got {resolved:?}");
1522 };
1523 assert_eq!(resolved_struct.repr(), TypeRepr::align(16));
1524 let resolved_fields = resolved_struct.get();
1525 assert_eq!(resolved_fields.fields()[0].name.as_deref(), Some("prefix"));
1526 assert_eq!(resolved_fields.fields()[1].name.as_deref(), Some("suffix"));
1527
1528 let TypeExpr::Struct(converted) = TypeExpr::from(resolved) else {
1529 panic!("expected concrete struct to convert back to struct type expression");
1530 };
1531 assert_eq!(*converted.repr, TypeRepr::align(16));
1532 assert_eq!(converted.fields[0].name.as_str(), "prefix");
1533 assert_eq!(converted.fields[1].name.as_str(), "suffix");
1534 }
1535}