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#[derive(Debug, Copy, Clone, Eq, PartialEq, Hash)]
12pub enum VariantType {
13 Struct,
23 Tuple,
31 Unit,
39}
40
41#[derive(Debug, Error)]
43pub enum VariantInfoError {
44 #[error("variant type mismatch: expected {expected:?}, received {received:?}")]
48 TypeMismatch {
49 expected: VariantType,
51 received: VariantType,
53 },
54}
55
56#[derive(Clone, Debug)]
58pub enum VariantInfo {
59 Struct(StructVariantInfo),
69 Tuple(TupleVariantInfo),
77 Unit(UnitVariantInfo),
85}
86
87impl VariantInfo {
88 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 #[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 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
145impl 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#[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 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 #[cfg(feature = "reflect_documentation")]
182 pub fn with_docs(self, docs: Option<&'static str>) -> Self {
183 Self { docs, ..self }
184 }
185
186 pub fn with_custom_attributes(self, custom_attributes: CustomAttributes) -> Self {
188 Self {
189 custom_attributes,
190 ..self
191 }
192 }
193
194 pub fn name(&self) -> &'static str {
196 self.name
197 }
198
199 pub fn field_names(&self) -> &[&'static str] {
201 &self.field_names
202 }
203
204 pub fn field(&self, name: &str) -> Option<&NamedField> {
206 self.field_indices
207 .get(name)
208 .map(|index| &self.fields[*index])
209 }
210
211 pub fn field_at(&self, index: usize) -> Option<&NamedField> {
213 self.fields.get(index)
214 }
215
216 pub fn index_of(&self, name: &str) -> Option<usize> {
218 self.field_indices.get(name).copied()
219 }
220
221 pub fn iter(&self) -> Iter<'_, NamedField> {
223 self.fields.iter()
224 }
225
226 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 #[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#[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 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 #[cfg(feature = "reflect_documentation")]
272 pub fn with_docs(self, docs: Option<&'static str>) -> Self {
273 Self { docs, ..self }
274 }
275
276 pub fn with_custom_attributes(self, custom_attributes: CustomAttributes) -> Self {
278 Self {
279 custom_attributes,
280 ..self
281 }
282 }
283
284 pub fn name(&self) -> &'static str {
286 self.name
287 }
288
289 pub fn field_at(&self, index: usize) -> Option<&UnnamedField> {
291 self.fields.get(index)
292 }
293
294 pub fn iter(&self) -> Iter<'_, UnnamedField> {
296 self.fields.iter()
297 }
298
299 pub fn field_len(&self) -> usize {
301 self.fields.len()
302 }
303
304 #[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#[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 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 #[cfg(feature = "reflect_documentation")]
335 pub fn with_docs(self, docs: Option<&'static str>) -> Self {
336 Self { docs, ..self }
337 }
338
339 pub fn with_custom_attributes(self, custom_attributes: CustomAttributes) -> Self {
341 Self {
342 custom_attributes,
343 ..self
344 }
345 }
346
347 pub fn name(&self) -> &'static str {
349 self.name
350 }
351
352 #[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}