Skip to main content

wasm_wave/value/
wit.rs

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
9/// Resolves a [`value::Type`] from the given [`wit_parser::Resolve`] and [`TypeId`].
10/// # Panics
11/// Panics if `type_id` is not valid in `resolve`.
12pub fn resolve_wit_type(resolve: &Resolve, type_id: TypeId) -> Result<value::Type, WasmValueError> {
13    TypeResolver { resolve }.resolve_type_id(type_id)
14}
15
16/// Resolves a [`value::FuncType`] from the given [`wit_parser::Resolve`] and [`Function`].
17/// # Panics
18/// Panics if `function`'s types are not valid in `resolve`.
19pub 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        // Recursively resolve any type defs.
62        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}