Skip to main content

bevy_reflect/enums/
variants.rs

1use crate::{
2    attributes::{impl_custom_attribute_methods, CustomAttributes},
3    NamedField, UnnamedField,
4};
5use alloc::boxed::Box;
6use bevy_platform::collections::HashMap;
7use core::slice::Iter;
8use thiserror::Error;
9
10/// Describes the form of an enum variant.
11#[derive(Debug, Copy, Clone, Eq, PartialEq, Hash)]
12pub enum VariantType {
13    /// Struct enums take the form:
14    ///
15    /// ```
16    /// enum MyEnum {
17    ///   A {
18    ///     foo: usize
19    ///   }
20    /// }
21    /// ```
22    Struct,
23    /// Tuple enums take the form:
24    ///
25    /// ```
26    /// enum MyEnum {
27    ///   A(usize)
28    /// }
29    /// ```
30    Tuple,
31    /// Unit enums take the form:
32    ///
33    /// ```
34    /// enum MyEnum {
35    ///   A
36    /// }
37    /// ```
38    Unit,
39}
40
41/// A [`VariantInfo`]-specific error.
42#[derive(Debug, Error)]
43pub enum VariantInfoError {
44    /// Caused when a variant was expected to be of a certain [type], but was not.
45    ///
46    /// [type]: VariantType
47    #[error("variant type mismatch: expected {expected:?}, received {received:?}")]
48    TypeMismatch {
49        /// Expected variant type.
50        expected: VariantType,
51        /// Received variant type.
52        received: VariantType,
53    },
54}
55
56/// A container for compile-time enum variant info.
57#[derive(Clone, Debug)]
58pub enum VariantInfo {
59    /// Struct enums take the form:
60    ///
61    /// ```
62    /// enum MyEnum {
63    ///   A {
64    ///     foo: usize
65    ///   }
66    /// }
67    /// ```
68    Struct(StructVariantInfo),
69    /// Tuple enums take the form:
70    ///
71    /// ```
72    /// enum MyEnum {
73    ///   A(usize)
74    /// }
75    /// ```
76    Tuple(TupleVariantInfo),
77    /// Unit enums take the form:
78    ///
79    /// ```
80    /// enum MyEnum {
81    ///   A
82    /// }
83    /// ```
84    Unit(UnitVariantInfo),
85}
86
87impl VariantInfo {
88    /// The name of the enum variant.
89    pub fn name(&self) -> &'static str {
90        match self {
91            Self::Struct(info) => info.name(),
92            Self::Tuple(info) => info.name(),
93            Self::Unit(info) => info.name(),
94        }
95    }
96
97    /// The docstring of the underlying variant, if any.
98    #[cfg(feature = "reflect_documentation")]
99    pub fn docs(&self) -> Option<&str> {
100        match self {
101            Self::Struct(info) => info.docs(),
102            Self::Tuple(info) => info.docs(),
103            Self::Unit(info) => info.docs(),
104        }
105    }
106
107    /// Returns the [type] of this variant.
108    ///
109    /// [type]: VariantType
110    pub fn variant_type(&self) -> VariantType {
111        match self {
112            Self::Struct(_) => VariantType::Struct,
113            Self::Tuple(_) => VariantType::Tuple,
114            Self::Unit(_) => VariantType::Unit,
115        }
116    }
117
118    impl_custom_attribute_methods!(
119        self,
120        match self {
121            Self::Struct(info) => info.custom_attributes(),
122            Self::Tuple(info) => info.custom_attributes(),
123            Self::Unit(info) => info.custom_attributes(),
124        },
125        "variant"
126    );
127}
128
129macro_rules! impl_cast_method {
130    ($name:ident : $kind:ident => $info:ident) => {
131        #[doc = concat!("Attempts a cast to [`", stringify!($info), "`].")]
132        #[doc = concat!("\n\nReturns an error if `self` is not [`VariantInfo::", stringify!($kind), "`].")]
133        pub fn $name(&self) -> Result<&$info, VariantInfoError> {
134            match self {
135                Self::$kind(info) => Ok(info),
136                _ => Err(VariantInfoError::TypeMismatch {
137                    expected: VariantType::$kind,
138                    received: self.variant_type(),
139                }),
140            }
141        }
142    };
143}
144
145/// Conversion convenience methods for [`VariantInfo`].
146impl VariantInfo {
147    impl_cast_method!(as_struct_variant: Struct => StructVariantInfo);
148    impl_cast_method!(as_tuple_variant: Tuple => TupleVariantInfo);
149    impl_cast_method!(as_unit_variant: Unit => UnitVariantInfo);
150}
151
152/// Type info for struct variants.
153#[derive(Clone, Debug)]
154pub struct StructVariantInfo {
155    name: &'static str,
156    fields: Box<[NamedField]>,
157    field_names: Box<[&'static str]>,
158    field_indices: HashMap<&'static str, usize>,
159    custom_attributes: CustomAttributes,
160    #[cfg(feature = "reflect_documentation")]
161    docs: Option<&'static str>,
162}
163
164impl StructVariantInfo {
165    /// Create a new [`StructVariantInfo`].
166    pub fn new(name: &'static str, fields: &[NamedField]) -> Self {
167        let field_indices = Self::collect_field_indices(fields);
168        let field_names = fields.iter().map(NamedField::name).collect();
169        Self {
170            name,
171            fields: fields.to_vec().into_boxed_slice(),
172            field_names,
173            field_indices,
174            custom_attributes: CustomAttributes::default(),
175            #[cfg(feature = "reflect_documentation")]
176            docs: None,
177        }
178    }
179
180    /// Sets the docstring for this variant.
181    #[cfg(feature = "reflect_documentation")]
182    pub fn with_docs(self, docs: Option<&'static str>) -> Self {
183        Self { docs, ..self }
184    }
185
186    /// Sets the custom attributes for this variant.
187    pub fn with_custom_attributes(self, custom_attributes: CustomAttributes) -> Self {
188        Self {
189            custom_attributes,
190            ..self
191        }
192    }
193
194    /// The name of this variant.
195    pub fn name(&self) -> &'static str {
196        self.name
197    }
198
199    /// A slice containing the names of all fields in order.
200    pub fn field_names(&self) -> &[&'static str] {
201        &self.field_names
202    }
203
204    /// Get the field with the given name.
205    pub fn field(&self, name: &str) -> Option<&NamedField> {
206        self.field_indices
207            .get(name)
208            .map(|index| &self.fields[*index])
209    }
210
211    /// Get the field at the given index.
212    pub fn field_at(&self, index: usize) -> Option<&NamedField> {
213        self.fields.get(index)
214    }
215
216    /// Get the index of the field with the given name.
217    pub fn index_of(&self, name: &str) -> Option<usize> {
218        self.field_indices.get(name).copied()
219    }
220
221    /// Iterate over the fields of this variant.
222    pub fn iter(&self) -> Iter<'_, NamedField> {
223        self.fields.iter()
224    }
225
226    /// The total number of fields in this variant.
227    pub fn field_len(&self) -> usize {
228        self.fields.len()
229    }
230
231    fn collect_field_indices(fields: &[NamedField]) -> HashMap<&'static str, usize> {
232        fields
233            .iter()
234            .enumerate()
235            .map(|(index, field)| (field.name(), index))
236            .collect()
237    }
238
239    /// The docstring of this variant, if any.
240    #[cfg(feature = "reflect_documentation")]
241    pub fn docs(&self) -> Option<&'static str> {
242        self.docs
243    }
244
245    impl_custom_attribute_methods!(self.custom_attributes, "variant");
246}
247
248/// Type info for tuple variants.
249#[derive(Clone, Debug)]
250pub struct TupleVariantInfo {
251    name: &'static str,
252    fields: Box<[UnnamedField]>,
253    custom_attributes: CustomAttributes,
254    #[cfg(feature = "reflect_documentation")]
255    docs: Option<&'static str>,
256}
257
258impl TupleVariantInfo {
259    /// Create a new [`TupleVariantInfo`].
260    pub fn new(name: &'static str, fields: &[UnnamedField]) -> Self {
261        Self {
262            name,
263            fields: fields.to_vec().into_boxed_slice(),
264            custom_attributes: CustomAttributes::default(),
265            #[cfg(feature = "reflect_documentation")]
266            docs: None,
267        }
268    }
269
270    /// Sets the docstring for this variant.
271    #[cfg(feature = "reflect_documentation")]
272    pub fn with_docs(self, docs: Option<&'static str>) -> Self {
273        Self { docs, ..self }
274    }
275
276    /// Sets the custom attributes for this variant.
277    pub fn with_custom_attributes(self, custom_attributes: CustomAttributes) -> Self {
278        Self {
279            custom_attributes,
280            ..self
281        }
282    }
283
284    /// The name of this variant.
285    pub fn name(&self) -> &'static str {
286        self.name
287    }
288
289    /// Get the field at the given index.
290    pub fn field_at(&self, index: usize) -> Option<&UnnamedField> {
291        self.fields.get(index)
292    }
293
294    /// Iterate over the fields of this variant.
295    pub fn iter(&self) -> Iter<'_, UnnamedField> {
296        self.fields.iter()
297    }
298
299    /// The total number of fields in this variant.
300    pub fn field_len(&self) -> usize {
301        self.fields.len()
302    }
303
304    /// The docstring of this variant, if any.
305    #[cfg(feature = "reflect_documentation")]
306    pub fn docs(&self) -> Option<&'static str> {
307        self.docs
308    }
309
310    impl_custom_attribute_methods!(self.custom_attributes, "variant");
311}
312
313/// Type info for unit variants.
314#[derive(Clone, Debug)]
315pub struct UnitVariantInfo {
316    name: &'static str,
317    custom_attributes: CustomAttributes,
318    #[cfg(feature = "reflect_documentation")]
319    docs: Option<&'static str>,
320}
321
322impl UnitVariantInfo {
323    /// Create a new [`UnitVariantInfo`].
324    pub fn new(name: &'static str) -> Self {
325        Self {
326            name,
327            custom_attributes: CustomAttributes::default(),
328            #[cfg(feature = "reflect_documentation")]
329            docs: None,
330        }
331    }
332
333    /// Sets the docstring for this variant.
334    #[cfg(feature = "reflect_documentation")]
335    pub fn with_docs(self, docs: Option<&'static str>) -> Self {
336        Self { docs, ..self }
337    }
338
339    /// Sets the custom attributes for this variant.
340    pub fn with_custom_attributes(self, custom_attributes: CustomAttributes) -> Self {
341        Self {
342            custom_attributes,
343            ..self
344        }
345    }
346
347    /// The name of this variant.
348    pub fn name(&self) -> &'static str {
349        self.name
350    }
351
352    /// The docstring of this variant, if any.
353    #[cfg(feature = "reflect_documentation")]
354    pub fn docs(&self) -> Option<&'static str> {
355        self.docs
356    }
357
358    impl_custom_attribute_methods!(self.custom_attributes, "variant");
359}
360
361#[cfg(test)]
362mod tests {
363    use super::*;
364    use crate::{Reflect, Typed};
365
366    #[test]
367    fn should_return_error_on_invalid_cast() {
368        #[derive(Reflect)]
369        enum Foo {
370            Bar,
371        }
372
373        let info = Foo::type_info().as_enum().unwrap();
374        let variant = info.variant_at(0).unwrap();
375        assert!(matches!(
376            variant.as_tuple_variant(),
377            Err(VariantInfoError::TypeMismatch {
378                expected: VariantType::Tuple,
379                received: VariantType::Unit
380            })
381        ));
382    }
383}