From 7311b7638de76ba68f197f891248b426e656e730 Mon Sep 17 00:00:00 2001 From: Michael Williamson Date: Fri, 3 Jul 2026 15:07:18 +0100 Subject: Add constructor functions for Rust enum variants --- examples/02-sum/output/rust/src/shapes.rs | 10 ++++ .../output/rust/src/arithmetic.rs | 10 ++++ .../output/rust/src/main.rs | 14 +++--- .../generators/rusttypes/RustTypesGenerator.java | 55 +++++++++++++++++++++- .../compiler/output/lang/rust/RustWriter.java | 2 +- .../output/lang/rust/ast/RustPathExprSegment.java | 8 ++++ .../output/lang/rust/ast/RustPathInExpression.java | 35 ++++++++++++++ .../compiler/output/lang/rust/RustWriterTests.java | 12 ++--- 8 files changed, 131 insertions(+), 15 deletions(-) diff --git a/examples/02-sum/output/rust/src/shapes.rs b/examples/02-sum/output/rust/src/shapes.rs index 958db0a..31c3db9 100644 --- a/examples/02-sum/output/rust/src/shapes.rs +++ b/examples/02-sum/output/rust/src/shapes.rs @@ -6,6 +6,16 @@ pub enum Shape { Triangle(crate::shapes::Triangle), } +impl Shape { + pub fn rectangle(value: crate::shapes::Rectangle) -> Self { + Self::Rectangle(value) + } + + pub fn triangle(value: crate::shapes::Triangle) -> Self { + Self::Triangle(value) + } +} + pub struct Rectangle { pub width: ::core::primitive::i32, pub height: ::core::primitive::i32, diff --git a/examples/03-inductive-data-types/output/rust/src/arithmetic.rs b/examples/03-inductive-data-types/output/rust/src/arithmetic.rs index a9679dc..d6944a9 100644 --- a/examples/03-inductive-data-types/output/rust/src/arithmetic.rs +++ b/examples/03-inductive-data-types/output/rust/src/arithmetic.rs @@ -5,6 +5,16 @@ pub enum Expression { Const(crate::arithmetic::Const), } +impl Expression { + pub fn add(value: crate::arithmetic::Add) -> Self { + Self::Add(::std::boxed::Box::new(value)) + } + + pub fn r#const(value: crate::arithmetic::Const) -> Self { + Self::Const(value) + } +} + pub struct Add { pub left: crate::arithmetic::Expression, pub right: crate::arithmetic::Expression, diff --git a/examples/03-inductive-data-types/output/rust/src/main.rs b/examples/03-inductive-data-types/output/rust/src/main.rs index 8c1f750..28f7cec 100644 --- a/examples/03-inductive-data-types/output/rust/src/main.rs +++ b/examples/03-inductive-data-types/output/rust/src/main.rs @@ -3,11 +3,11 @@ mod arithmetic; use arithmetic::{Add, Const, Expression}; fn main() { - let expression = Expression::Add(Box::new(Add { - left: Expression::Add(Box::new(Add { - left: Expression::Const(Const { value: 42 }), - right: Expression::Const(Const { value: 47 }), - })), - right: Expression::Const(Const { value: 52 }), - })); + let expression = Expression::add(Add { + left: Expression::add(Add { + left: Expression::r#const(Const { value: 42 }), + right: Expression::r#const(Const { value: 47 }), + }), + right: Expression::r#const(Const { value: 52 }), + }); } diff --git a/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttypes/RustTypesGenerator.java b/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttypes/RustTypesGenerator.java index 2e8b244..a7d39eb 100644 --- a/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttypes/RustTypesGenerator.java +++ b/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttypes/RustTypesGenerator.java @@ -171,7 +171,60 @@ public class RustTypesGenerator implements Generator { var rustEnum = new RustEnum(rustEnumName, rustVariants, sumDefinition.docComment()); - return List.of(rustEnum); + var rustCreateVariantFunctions = sumDefinition.variants().stream() + .map(variant -> { + var variantType = generateRustTypeExpression(variant.type(), context); + +// if (variant.isBox()) { +// variantType = RustTypePath.global("std", "boxed", "Box") +// .withArgs(List.of(variantType)); +// } +// + var variantName = generateVariantName(variant.type().value()); +// +// return new RustEnumVariant( +// variantName, +// new RustEnumVariantTuple(List.of( +// new RustTupleField(variantType) +// )) +// ); + + var innerValueName = new RustIdentifier("value"); + RustExpression innerValue = new RustPathInExpression(List.of( + new RustPathExprSegment( + new RustPathIdentSegmentIdentifier(innerValueName), + Optional.empty() + ) + )); + + if (variant.isBox()) { + innerValue = new RustCallExpression( + RustPathInExpression.global("std", "boxed", "Box", "new"), + List.of(innerValue) + ); + } + + return new RustFunction( + new RustIdentifier(Casing.upperCamelCaseToSnakeCase(variantName.value())), + List.of( + new RustFunctionParam(innerValueName, variantType) + ), + Optional.of(RustTypePath.selfType()), + Optional.of(new RustBlockExpression( + List.of(), + Optional.of( + new RustCallExpression( + RustPathInExpression.selfType(variantName), + List.of(innerValue) + ) + ) + )) + ); + }) + .toList(); + var rustImpl = new RustInherentImpl(RustTypePath.of(rustEnumName), rustCreateVariantFunctions); + + return List.of(rustEnum, rustImpl); } private RustType generateRustTypeExpression( diff --git a/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustWriter.java b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustWriter.java index 6db480e..b732390 100644 --- a/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustWriter.java +++ b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustWriter.java @@ -113,7 +113,7 @@ public class RustWriter implements AutoCloseable { } private void writeFunction(RustFunction function) throws IOException { - this.writer.write("fn "); + this.writer.write("pub fn "); this.writeIdentifier(function.name()); this.writer.write("("); diff --git a/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/ast/RustPathExprSegment.java b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/ast/RustPathExprSegment.java index 9e55ecf..7246ac0 100644 --- a/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/ast/RustPathExprSegment.java +++ b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/ast/RustPathExprSegment.java @@ -38,6 +38,14 @@ public record RustPathExprSegment(org.zwobble.hobgoblin.compiler.output.lang.rus } // Custom area start: org.zwobble.hobgoblin.compiler.output.lang.rust.ast.RustPathExprSegment body + public static RustPathExprSegment global() { + return new RustPathExprSegment(new RustPathIdentSegmentGlobal(), Optional.empty()); + } + + public static RustPathExprSegment selfType() { + return new RustPathExprSegment(new RustPathIdentSegmentSelfType(), Optional.empty()); + } + public RustPathExprSegment withArgs(List args) { if (this.args.isPresent()) { throw new IllegalArgumentException("Expression path segment already has args"); diff --git a/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/ast/RustPathInExpression.java b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/ast/RustPathInExpression.java index be28b7c..78baabc 100644 --- a/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/ast/RustPathInExpression.java +++ b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/ast/RustPathInExpression.java @@ -7,6 +7,7 @@ import java.util.ArrayList; import java.util.Arrays; import java.util.List; import java.util.Optional; +import java.util.stream.Stream; // Custom area end: org.zwobble.hobgoblin.compiler.output.lang.rust.ast.RustPathInExpression imports public record RustPathInExpression(java.util.List segments) implements org.zwobble.hobgoblin.compiler.output.lang.rust.ast.RustExpression { @@ -50,6 +51,40 @@ public record RustPathInExpression(java.util.List new RustPathExprSegment( + new RustPathIdentSegmentIdentifier(name), + Optional.empty() + )) + ) + .toList(); + return new RustPathInExpression(segments); + } + + private static RustPathInExpression qualified(RustPathExprSegment qualifier, String... names) { + var segments = Stream.concat( + Stream.of(qualifier), + Arrays.stream(names) + .map(name -> new RustPathExprSegment( + new RustPathIdentSegmentIdentifier(new RustIdentifier(name)), + Optional.empty() + )) + ) + .toList(); + return new RustPathInExpression(segments); + } + public RustPathInExpression withArgs(List args) { var segments = new ArrayList<>(this.segments); var lastSegment = segments.removeLast().withArgs(args); diff --git a/src/test/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustWriterTests.java b/src/test/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustWriterTests.java index 4c337d5..fe82fb9 100644 --- a/src/test/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustWriterTests.java +++ b/src/test/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustWriterTests.java @@ -136,7 +136,7 @@ public class RustWriterTests { var string = write(writer -> writer.writeAssociatedItem(rust)); assertThat(string, equalTo(""" - fn width();""")); + pub fn width();""")); } @Test @@ -150,7 +150,7 @@ public class RustWriterTests { var string = write(writer -> writer.writeAssociatedItem(rust)); assertThat(string, equalTo(""" - fn area(width: i32, height: u64);""")); + pub fn area(width: i32, height: u64);""")); } @Test @@ -163,7 +163,7 @@ public class RustWriterTests { var string = write(writer -> writer.writeAssociatedItem(rust)); assertThat(string, equalTo(""" - fn area() -> i32;""")); + pub fn area() -> i32;""")); } @Test @@ -178,7 +178,7 @@ public class RustWriterTests { var string = write(writer -> writer.writeAssociatedItem(rust)); assertThat(string, equalTo(""" - fn predicate() { + pub fn predicate() { false }""")); } @@ -206,9 +206,9 @@ public class RustWriterTests { assertThat(string, equalTo(""" impl Square { - fn width(); + pub fn width(); - fn height(); + pub fn height(); }""")); } -- cgit v1.2.3