1use alloc::{string::String, vec::Vec};
2use wit_parser::{
3 Enum, Flags, Function, Param, Record, Resolve, Result_, Tuple, Type, TypeDefKind, TypeId,
4 Variant,
5};
6
7use crate::{value, wasm::WasmValueError};
8
9pub fn resolve_wit_type(resolve: &Resolve, type_id: TypeId) -> Result<value::Type, WasmValueError> {
13 TypeResolver { resolve }.resolve_type_id(type_id)
14}
15
16pub fn resolve_wit_func_type(
20 resolve: &Resolve,
21 function: &Function,
22) -> Result<value::FuncType, WasmValueError> {
23 let resolver = TypeResolver { resolve };
24 let params = resolver.resolve_params(&function.params)?;
25 let results = match &function.result {
26 Some(ty) => [("".into(), resolver.resolve_type(*ty)?)].to_vec(),
27 None => Vec::new(),
28 };
29 value::FuncType::new(params, results)
30}
31
32struct TypeResolver<'a> {
33 resolve: &'a Resolve,
34}
35
36type ValueResult = Result<value::Type, WasmValueError>;
37
38impl<'a> TypeResolver<'a> {
39 fn resolve_type_id(&self, type_id: TypeId) -> ValueResult {
40 self.resolve(&self.resolve.types.get(type_id).unwrap().kind)
41 }
42
43 fn resolve_type(&self, ty: Type) -> ValueResult {
44 self.resolve(&TypeDefKind::Type(ty))
45 }
46
47 fn resolve_params(
48 &self,
49 params: &[Param],
50 ) -> Result<Vec<(String, value::Type)>, WasmValueError> {
51 params
52 .iter()
53 .map(|p| {
54 let ty = self.resolve_type(p.ty)?;
55 Ok((p.name.clone(), ty))
56 })
57 .collect()
58 }
59
60 fn resolve(&self, mut kind: &'a TypeDefKind) -> ValueResult {
61 while let &TypeDefKind::Type(Type::Id(id)) = kind {
63 kind = &self.resolve.types.get(id).unwrap().kind;
64 }
65
66 match kind {
67 TypeDefKind::Record(record) => self.resolve_record(record),
68 TypeDefKind::Flags(flags) => self.resolve_flags(flags),
69 TypeDefKind::Tuple(tuple) => self.resolve_tuple(tuple),
70 TypeDefKind::Variant(variant) => self.resolve_variant(variant),
71 TypeDefKind::Enum(enum_) => self.resolve_enum(enum_),
72 TypeDefKind::Option(some_type) => self.resolve_option(some_type),
73 TypeDefKind::Result(result) => self.resolve_result(result),
74 TypeDefKind::List(element_type) => self.resolve_list(element_type),
75 TypeDefKind::FixedLengthList(element_type, elements) => {
76 self.resolve_fixed_length_list(element_type, *elements)
77 }
78 TypeDefKind::Type(Type::Bool) => Ok(value::Type::BOOL),
79 TypeDefKind::Type(Type::U8) => Ok(value::Type::U8),
80 TypeDefKind::Type(Type::U16) => Ok(value::Type::U16),
81 TypeDefKind::Type(Type::U32) => Ok(value::Type::U32),
82 TypeDefKind::Type(Type::U64) => Ok(value::Type::U64),
83 TypeDefKind::Type(Type::S8) => Ok(value::Type::S8),
84 TypeDefKind::Type(Type::S16) => Ok(value::Type::S16),
85 TypeDefKind::Type(Type::S32) => Ok(value::Type::S32),
86 TypeDefKind::Type(Type::S64) => Ok(value::Type::S64),
87 TypeDefKind::Type(Type::F32) => Ok(value::Type::F32),
88 TypeDefKind::Type(Type::F64) => Ok(value::Type::F64),
89 TypeDefKind::Type(Type::Char) => Ok(value::Type::CHAR),
90 TypeDefKind::Type(Type::String) => Ok(value::Type::STRING),
91 TypeDefKind::Type(Type::Id(_)) => unreachable!(),
92 other => Err(WasmValueError::UnsupportedType(other.as_str().into())),
93 }
94 }
95
96 fn resolve_record(&self, record: &Record) -> ValueResult {
97 let fields = record
98 .fields
99 .iter()
100 .map(|f| Ok((f.name.as_str(), self.resolve_type(f.ty)?)))
101 .collect::<Result<Vec<_>, _>>()?;
102 Ok(value::Type::record(fields).unwrap())
103 }
104
105 fn resolve_flags(&self, flags: &Flags) -> ValueResult {
106 let names = flags.flags.iter().map(|f| f.name.as_str());
107 Ok(value::Type::flags(names).unwrap())
108 }
109
110 fn resolve_tuple(&self, tuple: &Tuple) -> ValueResult {
111 let types = tuple
112 .types
113 .iter()
114 .map(|ty| self.resolve_type(*ty))
115 .collect::<Result<Vec<_>, _>>()?;
116 Ok(value::Type::tuple(types).unwrap())
117 }
118
119 fn resolve_variant(&self, variant: &Variant) -> ValueResult {
120 let cases = variant
121 .cases
122 .iter()
123 .map(|case| {
124 Ok((
125 case.name.as_str(),
126 case.ty.map(|ty| self.resolve_type(ty)).transpose()?,
127 ))
128 })
129 .collect::<Result<Vec<_>, _>>()?;
130 Ok(value::Type::variant(cases).unwrap())
131 }
132
133 fn resolve_enum(&self, enum_: &Enum) -> ValueResult {
134 let cases = enum_.cases.iter().map(|c| c.name.as_str());
135 Ok(value::Type::enum_ty(cases).unwrap())
136 }
137
138 fn resolve_option(&self, some_type: &Type) -> ValueResult {
139 let some = self.resolve_type(*some_type)?;
140 Ok(value::Type::option(some))
141 }
142
143 fn resolve_result(&self, result: &Result_) -> ValueResult {
144 let ok = result.ok.map(|ty| self.resolve_type(ty)).transpose()?;
145 let err = result.err.map(|ty| self.resolve_type(ty)).transpose()?;
146 Ok(value::Type::result(ok, err))
147 }
148
149 fn resolve_list(&self, element_type: &Type) -> ValueResult {
150 let element_type = self.resolve_type(*element_type)?;
151 Ok(value::Type::list(element_type))
152 }
153
154 fn resolve_fixed_length_list(&self, element_type: &Type, elements: u32) -> ValueResult {
155 let element_type = self.resolve_type(*element_type)?;
156 Ok(value::Type::fixed_length_list(element_type, elements))
157 }
158}
159
160#[cfg(test)]
161mod tests {
162
163 use alloc::string::ToString;
164
165 use super::*;
166
167 #[test]
168 fn resolve_wit_type_smoke_test() {
169 let mut resolve = Resolve::new();
170 resolve
171 .push_str(
172 "test.wit",
173 "
174package test:types;
175interface types {
176 type uint8 = u8;
177}
178 ",
179 )
180 .unwrap();
181
182 let (type_id, _) = resolve.types.iter().next().unwrap();
183 let ty = resolve_wit_type(&resolve, type_id).unwrap();
184 assert_eq!(ty, value::Type::U8);
185 }
186
187 #[test]
188 fn resolve_wit_func_type_smoke_test() {
189 let mut resolve = Resolve::new();
190 resolve
191 .push_str(
192 "test.wit",
193 r#"
194package test:types;
195interface types {
196 type uint8 = u8;
197 no-results: func(a: uint8, b: string);
198 one-result: func(c: uint8, d: string) -> uint8;
199 named-results: func(e: uint8, f: string) -> tuple<u8, string>;
200}
201 "#,
202 )
203 .unwrap();
204
205 for (func_name, expected_display) in [
206 ("no-results", "func(a: u8, b: string)"),
207 ("one-result", "func(c: u8, d: string) -> u8"),
208 (
209 "named-results",
210 "func(e: u8, f: string) -> tuple<u8, string>",
211 ),
212 ] {
213 let function = resolve
214 .interfaces
215 .iter()
216 .flat_map(|(_, i)| &i.functions)
217 .find_map(|(name, function)| (name == func_name).then_some(function))
218 .unwrap();
219 let ty = resolve_wit_func_type(&resolve, function).unwrap();
220 assert_eq!(ty.to_string(), expected_display, "for {function:?}");
221 }
222 }
223}