Generate a Rust tagged union.
(
self,
union: Union,
parent_stack: Optional[List[Message]] = None,
)
| 783 | return lines |
| 784 | |
| 785 | def generate_union( |
| 786 | self, |
| 787 | union: Union, |
| 788 | parent_stack: Optional[List[Message]] = None, |
| 789 | ) -> List[str]: |
| 790 | """Generate a Rust tagged union.""" |
| 791 | lines: List[str] = [] |
| 792 | |
| 793 | union_name = self.get_type_identifier(union) |
| 794 | comment = self.format_type_id_comment(union, "//") |
| 795 | if comment: |
| 796 | lines.append(comment) |
| 797 | derives = ["::fory::ForyUnion"] |
| 798 | for trait in ("Clone", "Debug", "PartialEq", "Eq", "Hash"): |
| 799 | if self.union_supports_trait(union, trait, parent_stack): |
| 800 | derives.append(trait) |
| 801 | lines.append(f"#[derive({', '.join(derives)})]") |
| 802 | lines.append(f"pub enum {union_name} {{") |
| 803 | lines.append(" #[fory(unknown)]") |
| 804 | lines.append(" Unknown(::fory::UnknownCase),") |
| 805 | |
| 806 | for index, field in enumerate(union.fields): |
| 807 | variant_name = self.get_union_case_identifier(union, field) |
| 808 | pointer_type = self.get_field_pointer_type(field) |
| 809 | variant_type = self.generate_type( |
| 810 | field.field_type, |
| 811 | nullable=False, |
| 812 | ref=field.ref, |
| 813 | element_optional=field.element_optional, |
| 814 | element_ref=field.element_ref, |
| 815 | parent_stack=parent_stack, |
| 816 | pointer_type=pointer_type, |
| 817 | ) |
| 818 | variant_type = self.qualify_union_payload_type( |
| 819 | field.field_type, variant_type, variant_name |
| 820 | ) |
| 821 | variant_attrs = [f"id = {field.number}"] |
| 822 | if index == 0: |
| 823 | variant_attrs.append("default") |
| 824 | lines.append(f" #[fory({', '.join(variant_attrs)})]") |
| 825 | payload_attr = self.get_payload_field_attr(field) |
| 826 | if payload_attr: |
| 827 | lines.append( |
| 828 | f" {variant_name}(#[fory({payload_attr})] {variant_type})," |
| 829 | ) |
| 830 | else: |
| 831 | lines.append(f" {variant_name}({variant_type}),") |
| 832 | |
| 833 | lines.append("}") |
| 834 | lines.append("") |
| 835 | |
| 836 | if union.fields: |
| 837 | default_field = union.fields[0] |
| 838 | default_variant = self.get_union_case_identifier(union, default_field) |
| 839 | default_pointer_type = self.get_field_pointer_type(default_field) |
| 840 | default_type = self.generate_type( |
| 841 | default_field.field_type, |
| 842 | nullable=False, |
no test coverage detected