From c44f0e747f00ec71bd9c773bfea7d287e50ed77b Mon Sep 17 00:00:00 2001 From: Michael Williamson Date: Sun, 19 Jul 2026 21:41:55 +0100 Subject: Support sums in rust-transient-0 --- .../output/rust/src/gen/data/transient_0.rs | 30 +++++++++ examples/10-transient-0/output/rust/src/lib.rs | 17 ++++- .../rusttransient0/RustTransient0Generator.java | 78 +++++++++++++++++++++- 3 files changed, 123 insertions(+), 2 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 fbfc82d..0ea3f6f 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 @@ -170,6 +170,36 @@ pub fn decode_enum_with_variants(reader: &mut impl std::io::Read) -> ::std::io:: }) } +pub fn encode_sum_with_variants(value: &crate::data::SumWithVariants, writer: &mut impl std::io::Write) -> ::std::io::Result::<()> { + match value { + crate::data::SumWithVariants::VariantOne(value) => { + { + crate::transient_0::encode_int_32(&0, writer)?; + }; + { + crate::data::transient_0::encode_variant_one(value, writer)?; + }; + }, + crate::data::SumWithVariants::VariantTwo(value) => { + { + crate::transient_0::encode_int_32(&1, writer)?; + }; + { + crate::data::transient_0::encode_variant_two(value, writer)?; + }; + }, + }; + ::std::io::Result::Ok(()) +} + +pub fn decode_sum_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::SumWithVariants::VariantOne(crate::data::transient_0::decode_variant_one(reader)?), + 1 => crate::data::SumWithVariants::VariantTwo(crate::data::transient_0::decode_variant_two(reader)?), + _ => 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 c3da14c..3aa9751 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::{EnumWithVariants, InnerStruct, OuterStruct, StructWithBool, StructWithInt32, StructWithInt64, StructWithList, StructWithOption, StructWithString}; + use super::data::{EnumWithVariants, InnerStruct, OuterStruct, StructWithBool, StructWithInt32, StructWithInt64, StructWithList, StructWithOption, StructWithString, SumWithVariants, VariantOne, VariantTwo }; #[test] fn struct_with_bool() { @@ -128,6 +128,21 @@ mod test { ); } + #[test] + fn sum_with_variants() { + assert_round_trip_encoding( + SumWithVariants::VariantOne(VariantOne { a: 10 }), + super::data::transient_0::encode_sum_with_variants, + super::data::transient_0::decode_sum_with_variants, + ); + + assert_round_trip_encoding( + SumWithVariants::VariantTwo(VariantTwo { a: 25 }), + super::data::transient_0::encode_sum_with_variants, + super::data::transient_0::decode_sum_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/rusttransient0/RustTransient0Generator.java b/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttransient0/RustTransient0Generator.java index f2b8254..fd7bd34 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 @@ -406,7 +406,83 @@ public class RustTransient0Generator implements Generator { } case TypedSumDefinitionNode sumDefinition -> { - yield Stream.of(); + var encodeFunction = generateEncodeFunction( + sumDefinition.type(), + List.of( + new RustExpressionStatement(new RustMatchExpression( + RustPath.of(VALUE_NAME), + IntStream.range(0, sumDefinition.variants().size()) + .mapToObj(variantIndex -> { + var variant = sumDefinition.variants().get(variantIndex); + var variantType = variant.type().value(); + var variantPath = this.rustGenerator.generateRustTypeExpression(sumDefinition.type()) + .addSegment(RustPathSegment.of(this.rustGenerator.generateVariantName(variantType))); + + return new RustMatchArm( + new RustTupleStructPattern(variantPath, List.of(new RustIdentifierPattern(VALUE_NAME))), + new RustBlockExpression( + List.of( + new RustExpressionStatement(new RustBlockExpression( + generateEncode( + new RustPrefixExpression( + RustPrefixOperator.BORROW, + new RustIntegerLiteral(variantIndex, Optional.empty()) + ), + NativeTypes.INT_32 + ), + Optional.empty() + )), + new RustExpressionStatement(new RustBlockExpression( + generateEncode(RustPath.of(VALUE_NAME), variantType), + Optional.empty() + )) + ), + Optional.empty() + ) + ); + }) + .toList() + )) + ) + ); + + var decodeMatchArms = IntStream.range(0, sumDefinition.variants().size()) + .mapToObj(variantIndex -> { + var variant = sumDefinition.variants().get(variantIndex); + var variantType = variant.type().value(); + var variantPath = this.rustGenerator.generateRustTypeExpression(sumDefinition.type()) + .addSegment(RustPathSegment.of(this.rustGenerator.generateVariantName(variantType))); + + return new RustMatchArm( + new RustLiteralPattern(new RustIntegerLiteral(variantIndex, Optional.empty())), + new RustCallExpression(variantPath, List.of(generateDecode(variantType))) + ); + }) + .collect(toArrayList()); + + decodeMatchArms.add(new RustMatchArm( + new RustWildcardPattern(), + generateTodo() + )); + + var decodeFunction = generateDecodeFunction( + sumDefinition.type(), + new RustBlockExpression( + List.of(), + Optional.of(new RustMatchExpression( + new RustPrefixExpression( + RustPrefixOperator.BORROW, + generateDecode(NativeTypes.INT_32) + ), + decodeMatchArms + )) + ) + ); + + yield Stream.of( + encodeFunction, + decodeFunction + ); } }; } -- cgit v1.2.3