From 104245a9b01c379cca76c3114702415d7d296563 Mon Sep 17 00:00:00 2001 From: Michael Williamson Date: Sun, 19 Jul 2026 17:54:59 +0100 Subject: Support enums in rust-transient-0 --- .../output/rust/src/gen/data/transient_0.rs | 18 ++++++ examples/10-transient-0/output/rust/src/lib.rs | 23 +++++++- .../output/generators/rust/RustGenerator.java | 5 +- .../rusttransient0/RustTransient0Generator.java | 64 +++++++++++++++++++++- .../compiler/output/lang/rust/ast/RustPath.java | 2 +- .../hobgoblin/compiler/util/ListCollectors.java | 15 +++++ 6 files changed, 123 insertions(+), 4 deletions(-) create mode 100644 src/main/java/org/zwobble/hobgoblin/compiler/util/ListCollectors.java diff --git a/examples/10-transient-0/output/rust/src/gen/data/transient_0.rs b/examples/10-transient-0/output/rust/src/gen/data/transient_0.rs index c5bb458..fbfc82d 100644 --- a/examples/10-transient-0/output/rust/src/gen/data/transient_0.rs +++ b/examples/10-transient-0/output/rust/src/gen/data/transient_0.rs @@ -152,6 +152,24 @@ pub fn decode_struct_with_option(reader: &mut impl std::io::Read) -> ::std::io:: std::io::Result::Ok(crate::data::StructWithOption { a: a, b: b }) } +pub fn encode_enum_with_variants(value: &crate::data::EnumWithVariants, writer: &mut impl std::io::Write) -> ::std::io::Result::<()> { + crate::transient_0::encode_int_32(&match value { + crate::data::EnumWithVariants::Zero => 0, + crate::data::EnumWithVariants::One => 1, + crate::data::EnumWithVariants::Two => 2, + }, writer)?; + ::std::io::Result::Ok(()) +} + +pub fn decode_enum_with_variants(reader: &mut impl std::io::Read) -> ::std::io::Result:: { + std::io::Result::Ok(match &crate::transient_0::decode_int_32(reader)? { + 0 => crate::data::EnumWithVariants::Zero, + 1 => crate::data::EnumWithVariants::One, + 2 => crate::data::EnumWithVariants::Two, + _ => todo!(), + }) +} + pub fn encode_variant_one(value: &crate::data::VariantOne, writer: &mut impl std::io::Write) -> ::std::io::Result::<()> { crate::transient_0::encode_int_32(&value.a, writer)?; ::std::io::Result::Ok(()) diff --git a/examples/10-transient-0/output/rust/src/lib.rs b/examples/10-transient-0/output/rust/src/lib.rs index 3de2597..c3da14c 100644 --- a/examples/10-transient-0/output/rust/src/lib.rs +++ b/examples/10-transient-0/output/rust/src/lib.rs @@ -4,7 +4,7 @@ pub mod transient_0; #[cfg(test)] mod test { use std::io::Cursor; - use super::data::{InnerStruct, OuterStruct, StructWithBool, StructWithInt32, StructWithInt64, StructWithList, StructWithOption, StructWithString}; + use super::data::{EnumWithVariants, InnerStruct, OuterStruct, StructWithBool, StructWithInt32, StructWithInt64, StructWithList, StructWithOption, StructWithString}; #[test] fn struct_with_bool() { @@ -107,6 +107,27 @@ mod test { ); } + #[test] + fn enum_with_variants() { + assert_round_trip_encoding( + EnumWithVariants::Zero, + super::data::transient_0::encode_enum_with_variants, + super::data::transient_0::decode_enum_with_variants, + ); + + assert_round_trip_encoding( + EnumWithVariants::One, + super::data::transient_0::encode_enum_with_variants, + super::data::transient_0::decode_enum_with_variants, + ); + + assert_round_trip_encoding( + EnumWithVariants::Two, + super::data::transient_0::encode_enum_with_variants, + super::data::transient_0::decode_enum_with_variants, + ); + } + fn assert_round_trip_encoding( value: T, encode: impl Fn(&T, &mut Cursor>) -> std::io::Result<()>, diff --git a/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rust/RustGenerator.java b/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rust/RustGenerator.java index db572fb..0a0290d 100644 --- a/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rust/RustGenerator.java +++ b/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rust/RustGenerator.java @@ -46,7 +46,10 @@ public class RustGenerator { } case EnumType enumType -> { - throw new UnsupportedOperationException("TODO"); + yield generateTypePath( + enumType.namespaceName(), + enumType.name() + ); } case SimpleNativeType simpleNativeType -> { diff --git a/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttransient0/RustTransient0Generator.java b/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttransient0/RustTransient0Generator.java index 9dd4108..6295f17 100644 --- a/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttransient0/RustTransient0Generator.java +++ b/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttransient0/RustTransient0Generator.java @@ -16,8 +16,11 @@ import java.nio.file.Path; import java.util.ArrayList; import java.util.List; import java.util.Optional; +import java.util.stream.IntStream; import java.util.stream.Stream; +import static org.zwobble.hobgoblin.compiler.util.ListCollectors.toArrayList; + public class RustTransient0Generator implements Generator { public static final String NAME = "rust-transient-0"; @@ -295,7 +298,66 @@ public class RustTransient0Generator implements Generator { ) { return switch (statement) { case TypedEnumDefinitionNode enumDefinition -> { - yield Stream.of(); + var rustType = this.rustGenerator.generateRustTypeExpression(enumDefinition.type()); + + var encodeFunction = generateEncodeFunction( + enumDefinition.type(), + generateEncode(new RustPrefixExpression( + RustPrefixOperator.BORROW, + new RustMatchExpression( + RustPath.of(VALUE_NAME), + IntStream.range(0, enumDefinition.variants().size()) + .mapToObj(variantIndex -> { + var variant = enumDefinition.variants().get(variantIndex); + var variantPath = this.rustGenerator.generateRustTypeExpression(enumDefinition.type()) + .addSegment(RustPathSegment.of(this.rustGenerator.generateVariantName(variant.name()))); + + return new RustMatchArm( + new RustPathPattern(variantPath), + new RustIntegerLiteral(variantIndex, Optional.empty()) + ); + }) + .toList() + ) + ), NativeTypes.INT_32) + ); + + var decodeMatchArms = IntStream.range(0, enumDefinition.variants().size()) + .mapToObj(variantIndex -> { + var variant = enumDefinition.variants().get(variantIndex); + var variantPath = this.rustGenerator.generateRustTypeExpression(enumDefinition.type()) + .addSegment(RustPathSegment.of(this.rustGenerator.generateVariantName(variant.name()))); + + return new RustMatchArm( + new RustLiteralPattern(new RustIntegerLiteral(variantIndex, Optional.empty())), + variantPath + ); + }) + .collect(toArrayList()); + + decodeMatchArms.add(new RustMatchArm( + new RustWildcardPattern(), + generateTodo() + )); + + var decodeFunction = generateDecodeFunction( + enumDefinition.type(), + new RustBlockExpression( + List.of(), + Optional.of(new RustMatchExpression( + new RustPrefixExpression( + RustPrefixOperator.BORROW, + generateDecode(NativeTypes.INT_32) + ), + decodeMatchArms + )) + ) + ); + + yield Stream.of( + encodeFunction, + decodeFunction + ); } case TypedNativeTypeDefinitionNode nativeTypeDefinition -> { diff --git a/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/ast/RustPath.java b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/ast/RustPath.java index 8390293..5647001 100644 --- a/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/ast/RustPath.java +++ b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/ast/RustPath.java @@ -115,7 +115,7 @@ public record RustPath(java.util.List(this.segments); segments.add(segment); return new RustPath(segments); diff --git a/src/main/java/org/zwobble/hobgoblin/compiler/util/ListCollectors.java b/src/main/java/org/zwobble/hobgoblin/compiler/util/ListCollectors.java new file mode 100644 index 0000000..cf38abd --- /dev/null +++ b/src/main/java/org/zwobble/hobgoblin/compiler/util/ListCollectors.java @@ -0,0 +1,15 @@ +package org.zwobble.hobgoblin.compiler.util; + +import java.util.ArrayList; +import java.util.List; +import java.util.stream.Collector; + +public class ListCollectors { + public static Collector> toArrayList() { + return Collector.of( + ArrayList::new, + List::add, + (left, right) -> { left.addAll(right); return left; } + ); + } +} -- cgit v1.2.3