summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorMichael Williamson <mike@zwobble.org>2026-07-19 17:54:59 +0100
committerMichael Williamson <mike@zwobble.org>2026-07-19 17:54:59 +0100
commit104245a9b01c379cca76c3114702415d7d296563 (patch)
tree1bce77836f6eac64f39ec2ff3c78bfc1daa03b1c
parenta11c9415b63f45c316d3b08d656fa93eb9602f13 (diff)
Support enums in rust-transient-0
-rw-r--r--examples/10-transient-0/output/rust/src/gen/data/transient_0.rs18
-rw-r--r--examples/10-transient-0/output/rust/src/lib.rs23
-rw-r--r--src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rust/RustGenerator.java5
-rw-r--r--src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttransient0/RustTransient0Generator.java64
-rw-r--r--src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/ast/RustPath.java2
-rw-r--r--src/main/java/org/zwobble/hobgoblin/compiler/util/ListCollectors.java15
6 files changed, 123 insertions, 4 deletions
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::<crate::data::EnumWithVariants> {
+ 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<T: std::cmp::PartialEq + std::fmt::Debug>(
value: T,
encode: impl Fn(&T, &mut Cursor<Vec<u8>>) -> 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<org.zwobble.hobgoblin.compiler.output.lang
return new RustPath(segments);
}
- public RustExpression addSegment(RustPathSegment segment) {
+ public RustPath addSegment(RustPathSegment segment) {
var segments = new ArrayList<>(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 <T> Collector<T, ?, List<T>> toArrayList() {
+ return Collector.of(
+ ArrayList::new,
+ List::add,
+ (left, right) -> { left.addAll(right); return left; }
+ );
+ }
+}