Skip to main content

capnp/
schema.rs

1//! Convenience wrappers of the datatypes defined in schema.capnp.
2
3use crate::dynamic_value;
4use crate::introspect::{self, RawBrandedStructSchema, RawEnumSchema};
5use crate::private::layout;
6use crate::schema_capnp::{annotation, enumerant, field, node};
7use crate::struct_list;
8use crate::traits::{IndexMove, ListIter, ShortListIter};
9use crate::Result;
10
11/// A struct node, with generics applied.
12#[derive(Clone, Copy)]
13pub struct StructSchema {
14    pub(crate) raw: RawBrandedStructSchema,
15    pub(crate) proto: node::Reader<'static>,
16}
17
18impl StructSchema {
19    pub fn new(raw: RawBrandedStructSchema) -> Self {
20        let proto = crate::any_pointer::Reader::new(
21            layout::PointerReader::get_root_from_arena(raw.generic.arena).unwrap(),
22        )
23        .get_as()
24        .unwrap();
25        Self { raw, proto }
26    }
27
28    pub fn get_proto(&self) -> node::Reader<'static> {
29        self.proto
30    }
31
32    pub fn get_fields(self) -> crate::Result<FieldList> {
33        if let node::Struct(s) = self.proto.which()? {
34            Ok(FieldList {
35                fields: s.get_fields()?,
36                parent: self,
37            })
38        } else {
39            panic!()
40        }
41    }
42
43    pub fn get_field_by_discriminant(self, discriminant: u16) -> Result<Option<Field>> {
44        match self
45            .raw
46            .generic
47            .members_by_discriminant
48            .get(discriminant as usize)
49        {
50            None => Ok(None),
51            Some(&idx) => Ok(Some(self.get_fields()?.get(idx))),
52        }
53    }
54
55    /// Looks up a field by name using binary search. Returns `None` if no matching field is found.
56    pub fn find_field_by_name(&self, name: &str) -> Result<Option<Field>> {
57        let fields = self.get_fields()?;
58        let mut lower: usize = 0;
59        let mut upper: usize = self.raw.generic.members_by_name.len();
60
61        while lower < upper {
62            let mid: usize = (lower + upper) / 2;
63            let candidate_index = self.raw.generic.members_by_name[mid];
64            let candidate_name = fields.get(candidate_index).get_proto().get_name()?;
65
66            use core::cmp::Ordering;
67            match (&name).partial_cmp(&candidate_name) {
68                Some(Ordering::Equal) => return Ok(Some(fields.get(candidate_index))),
69                Some(Ordering::Greater) => lower = mid + 1,
70                Some(Ordering::Less) => upper = mid,
71                None => unreachable!(),
72            }
73        }
74        Ok(None)
75    }
76
77    /// Like `find_field_by_name()`, but returns an error if the field is not found.
78    pub fn get_field_by_name(&self, name: &str) -> Result<Field> {
79        if let Some(field) = self.find_field_by_name(name)? {
80            Ok(field)
81        } else {
82            let mut error = crate::Error::from_kind(crate::ErrorKind::FieldNotFound);
83            write!(error, "{name}");
84            Err(error)
85        }
86    }
87
88    pub fn get_union_fields(self) -> Result<FieldSubset> {
89        if let node::Struct(s) = self.proto.which()? {
90            Ok(FieldSubset {
91                fields: s.get_fields()?,
92                indices: self.raw.generic.members_by_discriminant,
93                parent: self,
94            })
95        } else {
96            panic!()
97        }
98    }
99
100    pub fn get_non_union_fields(self) -> Result<FieldSubset> {
101        if let node::Struct(s) = self.proto.which()? {
102            Ok(FieldSubset {
103                fields: s.get_fields()?,
104                indices: self.raw.generic.nonunion_members,
105                parent: self,
106            })
107        } else {
108            panic!()
109        }
110    }
111
112    pub fn get_annotations(self) -> Result<AnnotationList> {
113        Ok(AnnotationList {
114            annotations: self.proto.get_annotations()?,
115            child_index: None,
116            get_annotation_type: self.raw.annotation_types,
117        })
118    }
119}
120
121impl From<RawBrandedStructSchema> for StructSchema {
122    fn from(rs: RawBrandedStructSchema) -> StructSchema {
123        StructSchema::new(rs)
124    }
125}
126
127impl ::core::cmp::PartialEq for StructSchema {
128    fn eq(&self, other: &Self) -> bool {
129        self.raw == other.raw
130    }
131}
132
133impl ::core::cmp::Eq for StructSchema {}
134
135impl ::core::hash::Hash for StructSchema {
136    fn hash<H: ::core::hash::Hasher>(&self, state: &mut H) {
137        self.raw.hash(state);
138    }
139}
140
141impl ::core::fmt::Debug for StructSchema {
142    fn fmt(&self, f: &mut ::core::fmt::Formatter) -> ::core::fmt::Result {
143        // Two schemas with the same display name are unequal if their brandings
144        // differ, so also include the type id.
145        match self.proto.get_display_name().map(|n| n.to_str()) {
146            Ok(Ok(name)) => write!(f, "StructSchema({name}, {:?})", self.raw.type_id),
147            _ => write!(f, "StructSchema({:?})", self.raw),
148        }
149    }
150}
151
152/// A field of a struct, with generics applied.
153#[derive(Clone, Copy)]
154pub struct Field {
155    proto: field::Reader<'static>,
156    index: u16,
157    ty: introspect::Type,
158    pub(crate) parent: StructSchema,
159}
160
161impl Field {
162    pub fn get_proto(self) -> field::Reader<'static> {
163        self.proto
164    }
165
166    pub fn get_type(&self) -> introspect::Type {
167        self.ty
168    }
169
170    pub fn get_index(&self) -> u16 {
171        self.index
172    }
173
174    pub fn get_annotations(self) -> Result<AnnotationList> {
175        Ok(AnnotationList {
176            annotations: self.proto.get_annotations()?,
177            child_index: Some(self.index),
178            get_annotation_type: self.parent.raw.annotation_types,
179        })
180    }
181}
182
183impl ::core::cmp::PartialEq for Field {
184    fn eq(&self, other: &Self) -> bool {
185        self.parent == other.parent && self.index == other.index
186    }
187}
188impl ::core::cmp::Eq for Field {}
189impl ::core::hash::Hash for Field {
190    fn hash<H: ::core::hash::Hasher>(&self, state: &mut H) {
191        self.parent.hash(state);
192        self.index.hash(state);
193    }
194}
195
196impl ::core::fmt::Debug for Field {
197    fn fmt(&self, f: &mut ::core::fmt::Formatter) -> ::core::fmt::Result {
198        match self.proto.get_name().map(|n| n.to_str()) {
199            Ok(Ok(name)) => write!(f, "Field({name}, {:?})", self.parent),
200            _ => write!(f, "Field(index {}, {:?})", self.index, self.parent),
201        }
202    }
203}
204
205/// A list of fields of a struct, with generics applied.
206#[derive(Clone, Copy)]
207pub struct FieldList {
208    pub(crate) fields: crate::struct_list::Reader<'static, field::Owned>,
209    pub(crate) parent: StructSchema,
210}
211
212impl FieldList {
213    pub fn len(&self) -> u16 {
214        self.fields.len().try_into().unwrap()
215    }
216
217    pub fn is_empty(&self) -> bool {
218        self.len() == 0
219    }
220
221    pub fn get(self, index: u16) -> Field {
222        Field {
223            proto: self.fields.get(index as u32),
224            index,
225            ty: (self.parent.raw.field_types)(index),
226            parent: self.parent,
227        }
228    }
229
230    pub fn iter(self) -> ShortListIter<Self, Field> {
231        ShortListIter::new(self, self.len())
232    }
233}
234
235impl IndexMove<u16, Field> for FieldList {
236    fn index_move(&self, index: u16) -> Field {
237        self.get(index)
238    }
239}
240
241impl ::core::iter::IntoIterator for FieldList {
242    type Item = Field;
243    type IntoIter = ShortListIter<FieldList, Self::Item>;
244
245    fn into_iter(self) -> Self::IntoIter {
246        self.iter()
247    }
248}
249
250/// A list of a subset of fields of a struct, with generics applied.
251#[derive(Clone, Copy)]
252pub struct FieldSubset {
253    fields: struct_list::Reader<'static, field::Owned>,
254    indices: &'static [u16],
255    parent: StructSchema,
256}
257
258impl FieldSubset {
259    pub fn len(&self) -> u16 {
260        self.indices.len().try_into().unwrap()
261    }
262
263    pub fn is_empty(&self) -> bool {
264        self.len() == 0
265    }
266
267    pub fn get(self, index: u16) -> Field {
268        let index = self.indices[index as usize];
269        Field {
270            proto: self.fields.get(index as u32),
271            index,
272            ty: (self.parent.raw.field_types)(index),
273            parent: self.parent,
274        }
275    }
276
277    pub fn iter(self) -> ShortListIter<Self, Field> {
278        ShortListIter::new(self, self.len())
279    }
280}
281
282impl IndexMove<u16, Field> for FieldSubset {
283    fn index_move(&self, index: u16) -> Field {
284        self.get(index)
285    }
286}
287
288impl ::core::iter::IntoIterator for FieldSubset {
289    type Item = Field;
290    type IntoIter = ShortListIter<FieldSubset, Self::Item>;
291
292    fn into_iter(self) -> Self::IntoIter {
293        self.iter()
294    }
295}
296
297/// An enum, with generics applied. (Generics may affect types of annotations.)
298#[derive(Clone, Copy)]
299pub struct EnumSchema {
300    pub(crate) raw: RawEnumSchema,
301    pub(crate) proto: node::Reader<'static>,
302}
303
304impl EnumSchema {
305    pub fn new(raw: RawEnumSchema) -> Self {
306        let proto = crate::any_pointer::Reader::new(
307            layout::PointerReader::get_root_from_arena(raw.arena).unwrap(),
308        )
309        .get_as()
310        .unwrap();
311        Self { raw, proto }
312    }
313
314    pub fn get_proto(self) -> node::Reader<'static> {
315        self.proto
316    }
317
318    pub fn get_enumerants(self) -> crate::Result<EnumerantList> {
319        if let node::Enum(s) = self.proto.which()? {
320            Ok(EnumerantList {
321                enumerants: s.get_enumerants()?,
322                parent: self,
323            })
324        } else {
325            panic!()
326        }
327    }
328
329    pub fn get_annotations(self) -> Result<AnnotationList> {
330        Ok(AnnotationList {
331            annotations: self.proto.get_annotations()?,
332            child_index: None,
333            get_annotation_type: self.raw.annotation_types,
334        })
335    }
336}
337
338impl From<RawEnumSchema> for EnumSchema {
339    fn from(re: RawEnumSchema) -> EnumSchema {
340        EnumSchema::new(re)
341    }
342}
343
344impl ::core::cmp::PartialEq for EnumSchema {
345    fn eq(&self, other: &Self) -> bool {
346        self.raw == other.raw
347    }
348}
349
350impl ::core::cmp::Eq for EnumSchema {}
351
352impl ::core::hash::Hash for EnumSchema {
353    fn hash<H: ::core::hash::Hasher>(&self, state: &mut H) {
354        self.raw.hash(state);
355    }
356}
357
358impl ::core::fmt::Debug for EnumSchema {
359    fn fmt(&self, f: &mut ::core::fmt::Formatter) -> ::core::fmt::Result {
360        match self.proto.get_display_name().map(|n| n.to_str()) {
361            Ok(Ok(name)) => write!(f, "EnumSchema({name})"),
362            _ => write!(f, "EnumSchema({:?})", self.raw),
363        }
364    }
365}
366
367/// An enumerant, with generics applied. (Generics may affect types of annotations.)
368#[derive(Clone, Copy)]
369pub struct Enumerant {
370    ordinal: u16,
371    parent: EnumSchema,
372    proto: enumerant::Reader<'static>,
373}
374
375impl Enumerant {
376    pub fn get_containing_enum(self) -> EnumSchema {
377        self.parent
378    }
379
380    pub fn get_ordinal(self) -> u16 {
381        self.ordinal
382    }
383
384    pub fn get_proto(self) -> enumerant::Reader<'static> {
385        self.proto
386    }
387
388    pub fn get_annotations(self) -> Result<AnnotationList> {
389        Ok(AnnotationList {
390            annotations: self.proto.get_annotations()?,
391            child_index: Some(self.ordinal),
392            get_annotation_type: self.parent.raw.annotation_types,
393        })
394    }
395}
396
397impl ::core::cmp::PartialEq for Enumerant {
398    fn eq(&self, other: &Self) -> bool {
399        self.parent == other.parent && self.ordinal == other.ordinal
400    }
401}
402impl ::core::cmp::Eq for Enumerant {}
403impl ::core::hash::Hash for Enumerant {
404    fn hash<H: ::core::hash::Hasher>(&self, state: &mut H) {
405        self.parent.hash(state);
406        self.ordinal.hash(state);
407    }
408}
409
410impl ::core::fmt::Debug for Enumerant {
411    fn fmt(&self, f: &mut ::core::fmt::Formatter) -> ::core::fmt::Result {
412        match self.proto.get_name().map(|n| n.to_str()) {
413            Ok(Ok(name)) => write!(f, "Enumerant({name}, {:?})", self.parent),
414            _ => write!(f, "Enumerant(ordinal {}, {:?})", self.ordinal, self.parent),
415        }
416    }
417}
418
419/// A list of enumerants.
420#[derive(Clone, Copy)]
421pub struct EnumerantList {
422    enumerants: struct_list::Reader<'static, enumerant::Owned>,
423    parent: EnumSchema,
424}
425
426impl EnumerantList {
427    pub fn len(&self) -> u16 {
428        self.enumerants.len().try_into().unwrap()
429    }
430
431    pub fn is_empty(&self) -> bool {
432        self.len() == 0
433    }
434
435    pub fn get(self, ordinal: u16) -> Enumerant {
436        Enumerant {
437            proto: self.enumerants.get(ordinal as u32),
438            ordinal,
439            parent: self.parent,
440        }
441    }
442
443    pub fn iter(self) -> ShortListIter<Self, Enumerant> {
444        ShortListIter::new(self, self.len())
445    }
446}
447
448impl IndexMove<u16, Enumerant> for EnumerantList {
449    fn index_move(&self, index: u16) -> Enumerant {
450        self.get(index)
451    }
452}
453
454impl ::core::iter::IntoIterator for EnumerantList {
455    type Item = Enumerant;
456    type IntoIter = ShortListIter<Self, Self::Item>;
457
458    fn into_iter(self) -> Self::IntoIter {
459        self.iter()
460    }
461}
462
463/// An annotation.
464#[derive(Clone, Copy)]
465pub struct Annotation {
466    proto: annotation::Reader<'static>,
467    ty: introspect::Type,
468}
469
470impl Annotation {
471    /// Gets the value held in this annotation.
472    pub fn get_value(self) -> Result<dynamic_value::Reader<'static>> {
473        dynamic_value::Reader::new(self.proto.get_value()?, self.ty)
474    }
475
476    /// Gets the ID of the annotation node.
477    pub fn get_id(&self) -> u64 {
478        self.proto.get_id()
479    }
480
481    /// Gets the type of the value held in this annotation.
482    pub fn get_type(&self) -> introspect::Type {
483        self.ty
484    }
485}
486
487/// A list of annotations.
488#[derive(Clone, Copy)]
489pub struct AnnotationList {
490    annotations: struct_list::Reader<'static, annotation::Owned>,
491    child_index: Option<u16>,
492    get_annotation_type: fn(Option<u16>, u32) -> introspect::Type,
493}
494
495impl AnnotationList {
496    pub fn len(&self) -> u32 {
497        self.annotations.len()
498    }
499
500    pub fn is_empty(&self) -> bool {
501        self.len() == 0
502    }
503
504    pub fn get(self, index: u32) -> Annotation {
505        let proto = self.annotations.get(index);
506        let ty = (self.get_annotation_type)(self.child_index, index);
507        Annotation { proto, ty }
508    }
509
510    /// Returns the first annotation in the list that matches `id`.
511    /// Otherwise returns `None`.
512    pub fn find(self, id: u64) -> Option<Annotation> {
513        self.iter().find(|&annotation| annotation.get_id() == id)
514    }
515
516    pub fn iter(self) -> ListIter<Self, Annotation> {
517        ListIter::new(self, self.len())
518    }
519}
520
521impl IndexMove<u32, Annotation> for AnnotationList {
522    fn index_move(&self, index: u32) -> Annotation {
523        self.get(index)
524    }
525}
526
527impl ::core::iter::IntoIterator for AnnotationList {
528    type Item = Annotation;
529    type IntoIter = ListIter<Self, Self::Item>;
530
531    fn into_iter(self) -> Self::IntoIter {
532        self.iter()
533    }
534}
535
536#[cfg(test)]
537mod tests {
538    use crate::introspect::Introspect;
539
540    #[cfg(feature = "std")]
541    #[test]
542    fn fields_can_be_hashed() {
543        let crate::introspect::TypeVariant::Struct(struct_schema) =
544            crate::schema_capnp::node::Owned::introspect().which()
545        else {
546            panic!("Expected a struct schema");
547        };
548
549        let struct_schema = crate::schema::StructSchema::new(struct_schema);
550
551        let display_name = struct_schema.get_field_by_name("displayName").unwrap();
552        let id = struct_schema.get_field_by_name("id").unwrap();
553
554        let mut map = std::collections::HashMap::new();
555        map.insert(display_name, 1);
556        map.insert(id, 2);
557
558        assert_eq!(map.get(&display_name), Some(&1));
559        assert_eq!(map.get(&id), Some(&2));
560        assert_eq!(
561            map.get(&struct_schema.get_field_by_name("displayName").unwrap()),
562            Some(&1)
563        );
564        assert_eq!(
565            map.get(&struct_schema.get_field_by_name("id").unwrap()),
566            Some(&2)
567        );
568    }
569
570    #[test]
571    fn fields_can_be_compared() {
572        let crate::introspect::TypeVariant::Struct(struct_schema) =
573            crate::schema_capnp::node::Owned::introspect().which()
574        else {
575            panic!("Expected a struct schema");
576        };
577
578        let struct_schema = crate::schema::StructSchema::new(struct_schema);
579
580        let display_name = struct_schema.get_field_by_name("displayName").unwrap();
581        let id = struct_schema.get_field_by_name("id").unwrap();
582
583        assert_eq!(display_name, display_name);
584        assert_eq!(
585            display_name,
586            struct_schema.get_field_by_name("displayName").unwrap()
587        );
588        assert_eq!(id, id);
589        assert_eq!(id, struct_schema.get_field_by_name("id").unwrap());
590
591        assert_ne!(display_name, id);
592    }
593
594    #[cfg(feature = "std")]
595    #[test]
596    fn schemas_can_be_hashed() {
597        let node_schema = {
598            let crate::introspect::TypeVariant::Struct(schema) =
599                crate::schema_capnp::node::Owned::introspect().which()
600            else {
601                panic!("Expected a struct schema");
602            };
603
604            crate::schema::StructSchema::new(schema)
605        };
606        let cgr_schema = {
607            let crate::introspect::TypeVariant::Struct(schema) =
608                crate::schema_capnp::code_generator_request::Owned::introspect().which()
609            else {
610                panic!("Expected a struct schema");
611            };
612            crate::schema::StructSchema::new(schema)
613        };
614
615        let mut map = std::collections::HashMap::new();
616        map.insert(node_schema, 1);
617        map.insert(cgr_schema, 2);
618
619        assert_eq!(map.get(&node_schema), Some(&1));
620        assert_eq!(map.get(&cgr_schema), Some(&2));
621    }
622
623    #[test]
624    fn schemas_can_be_compared() {
625        let node_schema = {
626            let crate::introspect::TypeVariant::Struct(schema) =
627                crate::schema_capnp::node::Owned::introspect().which()
628            else {
629                panic!("Expected a struct schema");
630            };
631
632            crate::schema::StructSchema::new(schema)
633        };
634        let cgr_schema = {
635            let crate::introspect::TypeVariant::Struct(schema) =
636                crate::schema_capnp::code_generator_request::Owned::introspect().which()
637            else {
638                panic!("Expected a struct schema");
639            };
640            crate::schema::StructSchema::new(schema)
641        };
642
643        assert_eq!(node_schema, node_schema);
644        assert_eq!(cgr_schema, cgr_schema);
645        assert_ne!(node_schema, cgr_schema);
646    }
647
648    #[test]
649    fn enum_schemas_can_be_compared() {
650        let crate::introspect::TypeVariant::Enum(raw) =
651            crate::schema_capnp::ElementSize::introspect().which()
652        else {
653            panic!("Expected an enum schema");
654        };
655        let schema = crate::schema::EnumSchema::new(raw);
656
657        assert_eq!(schema, crate::schema::EnumSchema::new(raw));
658
659        let enumerants = schema.get_enumerants().unwrap();
660        assert_eq!(enumerants.get(0), enumerants.get(0));
661        assert_ne!(enumerants.get(0), enumerants.get(1));
662    }
663
664    #[cfg(feature = "std")]
665    #[test]
666    fn enumerants_can_be_hashed() {
667        let crate::introspect::TypeVariant::Enum(raw) =
668            crate::schema_capnp::ElementSize::introspect().which()
669        else {
670            panic!("Expected an enum schema");
671        };
672        let schema = crate::schema::EnumSchema::new(raw);
673        let enumerants = schema.get_enumerants().unwrap();
674
675        let mut map = std::collections::HashMap::new();
676        map.insert(enumerants.get(0), 0);
677        map.insert(enumerants.get(1), 1);
678
679        assert_eq!(map.get(&enumerants.get(0)), Some(&0));
680        assert_eq!(map.get(&enumerants.get(1)), Some(&1));
681    }
682
683    #[test]
684    fn type_variants_can_be_compared() {
685        use crate::introspect::TypeVariant;
686
687        assert_eq!(u32::introspect().which(), TypeVariant::UInt32);
688        assert_ne!(u32::introspect().which(), TypeVariant::Int32);
689        assert_eq!(
690            crate::schema_capnp::node::Owned::introspect().which(),
691            crate::schema_capnp::node::Owned::introspect().which()
692        );
693        assert_ne!(
694            crate::schema_capnp::node::Owned::introspect().which(),
695            crate::schema_capnp::code_generator_request::Owned::introspect().which()
696        );
697    }
698}