Skip to main content

submilli_engine/
type_info.rs

1use std::collections::{BTreeMap, BTreeSet};
2
3use serde::{Deserialize, Serialize};
4
5use crate::types::LiteralF64;
6use crate::{Shape, Type, TypedAst};
7
8#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
9pub struct TypeInfoId(pub u32);
10
11impl TypeInfoId {
12    pub fn as_u32(self) -> u32 {
13        self.0
14    }
15
16    pub fn as_usize(self) -> usize {
17        self.0 as usize
18    }
19}
20
21#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
22pub struct TypeInfoTable {
23    pub package_name: String,
24    pub types: Vec<TypeInfo>,
25}
26
27#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
28pub struct TypeInfo {
29    pub id: TypeInfoId,
30    pub kind: TypeInfoKind,
31}
32
33#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
34pub enum TypeInfoKind {
35    Null,
36    Undefined,
37    Boolean,
38    Number,
39    String,
40    NumberLiteral(LiteralF64),
41    StringLiteral(String),
42    BooleanLiteral(bool),
43    Array {
44        element: TypeInfoId,
45    },
46    Tuple {
47        elements: Vec<TypeInfoId>,
48        optional: usize,
49    },
50    Object {
51        fields: Vec<FieldInfo>,
52    },
53    Union {
54        members: Vec<TypeInfoId>,
55    },
56    Unsupported {
57        label: String,
58    },
59}
60
61#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
62pub struct FieldInfo {
63    pub name: String,
64    pub optional: bool,
65    pub type_id: TypeInfoId,
66}
67
68impl TypeInfoTable {
69    pub fn collect(package_name: impl Into<String>, ast: &TypedAst) -> Self {
70        Self::collect_from_shapes(package_name, &ast.shapes)
71    }
72
73    pub fn collect_from_shapes(package_name: impl Into<String>, shapes: &[Shape]) -> Self {
74        Self::collect_from_shapes_and_types(package_name, shapes, &[])
75    }
76
77    pub fn collect_from_shapes_and_types(
78        package_name: impl Into<String>,
79        shapes: &[Shape],
80        types: &[Type],
81    ) -> Self {
82        Self::collect_indexed(package_name, shapes, types).0
83    }
84
85    /// [`collect_from_shapes_and_types`](Self::collect_from_shapes_and_types),
86    /// plus the id each collected type was interned under.
87    pub fn collect_indexed(
88        package_name: impl Into<String>,
89        shapes: &[Shape],
90        types: &[Type],
91    ) -> (Self, TypeInfoIndex) {
92        let mut builder = TypeInfoBuilder::new(package_name.into());
93        for shape in shapes {
94            builder.intern_type(&canonical_type(shape));
95        }
96        for ty in types {
97            builder.intern_type(ty);
98        }
99        builder.finish()
100    }
101
102    pub fn is_empty(&self) -> bool {
103        self.types.is_empty()
104    }
105
106    pub fn get(&self, id: TypeInfoId) -> Option<&TypeInfo> {
107        self.types.get(id.as_usize()).filter(|info| info.id == id)
108    }
109
110    pub fn object_type_id(&self, ty: &Type) -> Option<TypeInfoId> {
111        let Type::Object { .. } = ty.peel() else {
112            return None;
113        };
114        self.type_id(ty).filter(|id| {
115            matches!(
116                self.get(*id).map(|info| &info.kind),
117                Some(TypeInfoKind::Object { .. })
118            )
119        })
120    }
121
122    pub fn type_id(&self, ty: &Type) -> Option<TypeInfoId> {
123        self.types.iter().find_map(|info| {
124            let mut seen = BTreeSet::new();
125            self.type_info_matches_type(info.id, ty, &mut seen)
126                .then_some(info.id)
127        })
128    }
129
130    pub fn supports_host_json_object(&self, ty: &Type) -> bool {
131        let Type::Object {
132            fields,
133            index: None,
134        } = ty.peel()
135        else {
136            return false;
137        };
138        fields
139            .values()
140            .all(|field| self.supports_host_json_type(&field.ty))
141    }
142
143    pub fn supports_host_json_type(&self, ty: &Type) -> bool {
144        match ty.peel() {
145            Type::Null
146            | Type::Undefined
147            | Type::Boolean
148            | Type::BooleanLiteral(_)
149            | Type::Number
150            | Type::NumberLiteral(_)
151            | Type::String
152            | Type::StringLiteral(_) => true,
153            Type::Object { .. } => self.supports_host_json_object(ty),
154            Type::Array(element) => self.supports_host_json_type(element),
155            Type::Tuple(elements) => elements
156                .iter()
157                .all(|element| self.supports_host_json_type(element)),
158            Type::Union(members)
159                if members
160                    .iter()
161                    .any(|m| matches!(m.peel(), Type::Null | Type::Undefined)) =>
162            {
163                members
164                    .iter()
165                    .filter(|m| !matches!(m.peel(), Type::Null | Type::Undefined))
166                    .count()
167                    == 1
168                    && members
169                        .iter()
170                        .find(|m| !matches!(m.peel(), Type::Null | Type::Undefined))
171                        .is_some_and(|m| self.supports_host_json_type(m))
172            }
173            _ => false,
174        }
175    }
176
177    fn type_info_matches_type(
178        &self,
179        id: TypeInfoId,
180        ty: &Type,
181        seen: &mut BTreeSet<(TypeInfoId, Type)>,
182    ) -> bool {
183        if !seen.insert((id, ty.clone())) {
184            return true;
185        }
186        let Some(info) = self.get(id) else {
187            return false;
188        };
189        match (&info.kind, ty.peel()) {
190            (TypeInfoKind::Null, Type::Null)
191            | (TypeInfoKind::Undefined, Type::Undefined)
192            | (TypeInfoKind::Boolean, Type::Boolean)
193            | (TypeInfoKind::Number, Type::Number)
194            | (TypeInfoKind::String, Type::String) => true,
195            (TypeInfoKind::NumberLiteral(a), Type::NumberLiteral(b)) => a == b,
196            (TypeInfoKind::StringLiteral(a), Type::StringLiteral(b)) => a == b,
197            (TypeInfoKind::BooleanLiteral(a), Type::BooleanLiteral(b)) => a == b,
198            (TypeInfoKind::Array { element }, Type::Array(expected)) => {
199                self.type_info_matches_type(*element, expected, seen)
200            }
201            (TypeInfoKind::Tuple { elements, optional }, Type::Tuple(expected)) => {
202                *optional == expected.optional
203                    && elements.len() == expected.len()
204                    && elements
205                        .iter()
206                        .zip(expected)
207                        .all(|(id, ty)| self.type_info_matches_type(*id, ty, seen))
208            }
209            (
210                TypeInfoKind::Object { fields },
211                Type::Object {
212                    fields: expected, ..
213                },
214            ) => {
215                fields.len() == expected.len()
216                    && fields.iter().zip(expected).all(|(info, (name, field))| {
217                        info.name == *name
218                            && info.optional == field.optional
219                            && self.type_info_matches_type(info.type_id, &field.ty, seen)
220                    })
221            }
222            (TypeInfoKind::Union { members }, Type::Union(expected)) => {
223                members.len() == expected.len()
224                    && members
225                        .iter()
226                        .zip(expected)
227                        .all(|(id, ty)| self.type_info_matches_type(*id, ty, seen))
228            }
229            _ => false,
230        }
231    }
232}
233
234/// The id each type was interned under when a [`TypeInfoTable`] was collected.
235/// [`TypeInfoTable::type_id`] finds a type by comparing it with every entry in
236/// turn, which is quadratic over a program's types; a collected type is found
237/// here directly.
238#[derive(Debug, Default)]
239pub struct TypeInfoIndex {
240    by_type: BTreeMap<Type, TypeInfoId>,
241}
242
243impl TypeInfoIndex {
244    /// [`TypeInfoTable::object_type_id`], looking a collected type up directly.
245    pub fn object_type_id(&self, table: &TypeInfoTable, ty: &Type) -> Option<TypeInfoId> {
246        let Some(id) = self.by_type.get(ty) else {
247            return table.object_type_id(ty);
248        };
249        matches!(
250            table.get(*id).map(|info| &info.kind),
251            Some(TypeInfoKind::Object { .. })
252        )
253        .then_some(*id)
254    }
255}
256
257struct TypeInfoBuilder {
258    table: TypeInfoTable,
259    by_type: BTreeMap<Type, TypeInfoId>,
260}
261
262impl TypeInfoBuilder {
263    fn new(package_name: String) -> Self {
264        Self {
265            table: TypeInfoTable {
266                package_name,
267                types: Vec::new(),
268            },
269            by_type: BTreeMap::new(),
270        }
271    }
272
273    fn finish(self) -> (TypeInfoTable, TypeInfoIndex) {
274        (
275            self.table,
276            TypeInfoIndex {
277                by_type: self.by_type,
278            },
279        )
280    }
281
282    fn intern_type(&mut self, ty: &Type) -> TypeInfoId {
283        if let Some(id) = self.by_type.get(ty) {
284            return *id;
285        }
286
287        let id = TypeInfoId(self.table.types.len() as u32);
288        self.by_type.insert(ty.clone(), id);
289        self.table.types.push(TypeInfo {
290            id,
291            kind: TypeInfoKind::Unsupported {
292                label: ty.to_string(),
293            },
294        });
295
296        let kind = self.kind_for_type(ty);
297        self.table.types[id.as_usize()].kind = kind;
298        id
299    }
300
301    fn kind_for_type(&mut self, ty: &Type) -> TypeInfoKind {
302        match ty.peel() {
303            Type::Null => TypeInfoKind::Null,
304            Type::Undefined => TypeInfoKind::Undefined,
305            Type::Boolean => TypeInfoKind::Boolean,
306            Type::Number => TypeInfoKind::Number,
307            Type::String => TypeInfoKind::String,
308            Type::NumberLiteral(n) => TypeInfoKind::NumberLiteral(*n),
309            Type::StringLiteral(s) => TypeInfoKind::StringLiteral(s.clone()),
310            Type::BooleanLiteral(b) => TypeInfoKind::BooleanLiteral(*b),
311            Type::Array(element) => TypeInfoKind::Array {
312                element: self.intern_type(element),
313            },
314            Type::Tuple(elements) => TypeInfoKind::Tuple {
315                optional: elements.optional,
316                elements: elements
317                    .iter()
318                    .map(|element| self.intern_type(element))
319                    .collect(),
320            },
321            Type::Object { fields, .. } => TypeInfoKind::Object {
322                fields: fields
323                    .iter()
324                    .map(|(name, field)| FieldInfo {
325                        name: name.clone(),
326                        optional: field.optional,
327                        type_id: self.intern_type(&field.ty),
328                    })
329                    .collect(),
330            },
331            Type::Union(members) => TypeInfoKind::Union {
332                members: members
333                    .iter()
334                    .map(|member| self.intern_type(member))
335                    .collect(),
336            },
337            other => TypeInfoKind::Unsupported {
338                label: other.to_string(),
339            },
340        }
341    }
342}
343
344fn canonical_type(shape: &Shape) -> Type {
345    match shape {
346        Shape::Object { fields, index } => Type::Object {
347            index: index.clone(),
348            fields: fields.clone(),
349        },
350        Shape::Array(elem) => Type::Array(elem.clone()),
351        Shape::Tuple(elements) => Type::Tuple(elements.clone()),
352        Shape::Union(members) => Type::Union(members.clone()),
353    }
354}
355
356#[cfg(test)]
357mod tests {
358    use std::collections::BTreeMap;
359
360    use crate::{ObjectField, Shape, Type, TypeInfoKind, TypeInfoTable};
361
362    #[test]
363    fn index_finds_the_same_object_type_id_as_the_scan() {
364        let inner = BTreeMap::from([("v".to_string(), ObjectField::required(Type::Number))]);
365        let inner_ty = Type::Object {
366            index: None,
367            fields: inner.clone(),
368        };
369        let outer =
370            BTreeMap::from([("inner".to_string(), ObjectField::required(inner_ty.clone()))]);
371        let outer_ty = Type::Object {
372            index: None,
373            fields: outer.clone(),
374        };
375        let (table, index) = TypeInfoTable::collect_indexed(
376            "main",
377            &[
378                Shape::Object {
379                    index: None,
380                    fields: outer,
381                },
382                Shape::Array(Box::new(Type::String)),
383            ],
384            std::slice::from_ref(&inner_ty),
385        );
386        for ty in [&outer_ty, &inner_ty] {
387            let id = index.object_type_id(&table, ty);
388            assert!(id.is_some(), "{ty} has type info");
389            assert_eq!(id, table.object_type_id(ty));
390        }
391        let array = Type::Array(Box::new(Type::String));
392        assert_eq!(index.object_type_id(&table, &array), None);
393        // A type the table never collected falls back to the structural scan.
394        let uncollected = Type::Object {
395            index: None,
396            fields: BTreeMap::from([("w".to_string(), ObjectField::required(Type::Number))]),
397        };
398        assert_eq!(index.object_type_id(&table, &uncollected), None);
399    }
400
401    #[test]
402    fn object_type_id_finds_collected_object_shape() {
403        let fields = BTreeMap::from([
404            ("id".to_string(), ObjectField::required(Type::Number)),
405            ("name".to_string(), ObjectField::required(Type::String)),
406        ]);
407        let ty = Type::Object {
408            index: None,
409            fields: fields.clone(),
410        };
411        let table = TypeInfoTable::collect_from_shapes(
412            "main",
413            &[Shape::Object {
414                index: None,
415                fields,
416            }],
417        );
418        let id = table
419            .object_type_id(&ty)
420            .expect("collected object shape should have TypeInfo");
421
422        assert!(matches!(
423            table.get(id).map(|info| &info.kind),
424            Some(TypeInfoKind::Object { fields }) if fields.len() == 2
425        ));
426    }
427}