summaryrefslogtreecommitdiff
path: root/src
diff options
context:
space:
mode:
authorMichael Williamson <mike@zwobble.org>2026-07-18 10:26:25 +0100
committerMichael Williamson <mike@zwobble.org>2026-07-18 10:26:25 +0100
commit25851fc06ab5b2271fb12ee35c5c181c1f9258aa (patch)
tree9515cc4a956834db59cb287189384dfe9bd9e591 /src
parent3ea197f9f1deec604e99f2bf2292acd3fc7d3ac7 (diff)
Implement precedence in Rust writer
Diffstat (limited to 'src')
-rw-r--r--src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustAssociativity.java7
-rw-r--r--src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustPrecedence.java27
-rw-r--r--src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustWriter.java103
-rw-r--r--src/test/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustWriterTests.java97
4 files changed, 190 insertions, 44 deletions
diff --git a/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustAssociativity.java b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustAssociativity.java
new file mode 100644
index 0000000..632e967
--- /dev/null
+++ b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustAssociativity.java
@@ -0,0 +1,7 @@
+package org.zwobble.hobgoblin.compiler.output.lang.rust;
+
+public enum RustAssociativity {
+ LEFT,
+ RIGHT,
+ NONE;
+}
diff --git a/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustPrecedence.java b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustPrecedence.java
new file mode 100644
index 0000000..c3038ca
--- /dev/null
+++ b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustPrecedence.java
@@ -0,0 +1,27 @@
+package org.zwobble.hobgoblin.compiler.output.lang.rust;
+
+public enum RustPrecedence {
+ PRIMARY,
+ PATHS,
+ METHOD_CALLS,
+ FIELD_EXPRESSIONS,
+ FUNCTION_CALLS,
+ TRY_PROPAGATION,
+ PREFIX,
+ AS,
+ MULTIPLICATIVE,
+ ADDITIVE,
+ SHIFT,
+ BITWISE_AND,
+ BITWISE_XOR,
+ BITWISE_OR,
+ COMPARISON,
+ LOGICAL_AND,
+ LOGICAL_OR,
+ RANGE,
+ ASSIGN;
+
+ public int value() {
+ return -this.ordinal();
+ }
+}
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 e4b24a8..b584c69 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
@@ -294,7 +294,7 @@ public class RustWriter implements AutoCloseable {
}
private void writeExpressionStatement(RustExpressionStatement expressionStatement) throws IOException {
- this.writeExpression(expressionStatement.expression());
+ this.writeTopLevelExpression(expressionStatement.expression());
this.writer.write(";");
}
@@ -307,13 +307,13 @@ public class RustWriter implements AutoCloseable {
this.writeIdentifier(letStatement.variableName());
this.writer.write(" = ");
- this.writeExpression(letStatement.value());
+ this.writeTopLevelExpression(letStatement.value());
this.writer.write(";");
}
// == Expressions ==
- void writeExpression(RustExpression expression) throws IOException {
+ void writeTopLevelExpression(RustExpression expression) throws IOException {
switch (expression) {
case RustArrayRepeatExpression arrayRepeatExpression -> {
this.writeArrayRepeatExpression(arrayRepeatExpression);
@@ -369,16 +369,69 @@ public class RustWriter implements AutoCloseable {
}
}
+ void writeSubExpression(
+ RustExpression expression,
+ RustPrecedence parentPrecedence,
+ boolean matchesAssociativity
+ ) throws IOException {
+ var precedence = precedence(expression);
+ var requiresParens = (
+ parentPrecedence.value() > precedence.value() ||
+ (parentPrecedence.value() == precedence.value() && !matchesAssociativity)
+ );
+
+ if (requiresParens) {
+ this.writer.write("(");
+ }
+
+ this.writeTopLevelExpression(expression);
+
+ if (requiresParens) {
+ this.writer.write(")");
+ }
+ }
+
+ private RustPrecedence precedence(RustExpression expression) {
+ return switch (expression) {
+ case RustArrayRepeatExpression _ -> RustPrecedence.PRIMARY;
+ case RustBinaryExpression binaryExpression -> precedence(binaryExpression.operator());
+ case RustBlockExpression _ -> RustPrecedence.PRIMARY;
+ case RustBoolLiteral _ -> RustPrecedence.PRIMARY;
+ case RustCallExpression _ -> RustPrecedence.FUNCTION_CALLS;
+ case RustFieldExpression _ -> RustPrecedence.FIELD_EXPRESSIONS;
+ case RustIfExpression _ -> RustPrecedence.PRIMARY;
+ case RustIntegerLiteral _ -> RustPrecedence.PRIMARY;
+ case RustPath _ -> RustPrecedence.PATHS;
+ case RustPrefixExpression _ -> RustPrecedence.PREFIX;
+ case RustStructExpression _ -> RustPrecedence.PRIMARY;
+ case RustTryPropagationExpression _ -> RustPrecedence.TRY_PROPAGATION;
+ case RustTupleExpression _ -> RustPrecedence.PRIMARY;
+ };
+ }
+
+ private RustPrecedence precedence(RustBinaryOperator operator) {
+ return switch (operator) {
+ case ADD -> RustPrecedence.ADDITIVE;
+ case EQUAL -> RustPrecedence.COMPARISON;
+ case NOT_EQUAL -> RustPrecedence.COMPARISON;
+ };
+ }
+
private void writeArrayRepeatExpression(RustArrayRepeatExpression arrayRepeatExpression) throws IOException {
this.writer.write("[");
- this.writeExpression(arrayRepeatExpression.repeatOperand());
+ this.writeTopLevelExpression(arrayRepeatExpression.repeatOperand());
this.writer.write("; ");
- this.writeExpression(arrayRepeatExpression.lengthOperand());
+ this.writeTopLevelExpression(arrayRepeatExpression.lengthOperand());
this.writer.write("]");
}
private void writeBinaryExpression(RustBinaryExpression binaryExpression) throws IOException {
- this.writeExpression(binaryExpression.left());
+ var associativity = associativity(binaryExpression);
+ this.writeSubExpression(
+ binaryExpression.left(),
+ precedence(binaryExpression),
+ associativity == RustAssociativity.LEFT
+ );
this.writer.write(" ");
this.writer.write(switch (binaryExpression.operator()) {
case ADD -> "+";
@@ -386,7 +439,19 @@ public class RustWriter implements AutoCloseable {
case NOT_EQUAL -> "!=";
});
this.writer.write(" ");
- this.writeExpression(binaryExpression.right());
+ this.writeSubExpression(
+ binaryExpression.right(),
+ precedence(binaryExpression),
+ associativity == RustAssociativity.RIGHT
+ );
+ }
+
+ private RustAssociativity associativity(RustBinaryExpression binaryExpression) {
+ return switch (binaryExpression.operator()) {
+ case ADD -> RustAssociativity.LEFT;
+ case EQUAL -> RustAssociativity.NONE;
+ case NOT_EQUAL -> RustAssociativity.NONE;
+ };
}
private void writeBlockExpression(RustBlockExpression blockExpression) throws IOException {
@@ -400,7 +465,7 @@ public class RustWriter implements AutoCloseable {
if (blockExpression.finalOperand().isPresent()) {
this.writer.newLine();
- this.writeExpression(blockExpression.finalOperand().get());
+ this.writeTopLevelExpression(blockExpression.finalOperand().get());
}
this.writer.dedent();
@@ -413,11 +478,11 @@ public class RustWriter implements AutoCloseable {
}
private void writeCallExpression(RustCallExpression callExpression) throws IOException {
- this.writeExpression(callExpression.function());
+ this.writeSubExpression(callExpression.function(), precedence(callExpression), true);
this.writer.write("(");
writeWithSeparator(
callExpression.args(),
- this::writeExpression,
+ this::writeTopLevelExpression,
() -> {
this.writer.write(", ");
}
@@ -426,14 +491,14 @@ public class RustWriter implements AutoCloseable {
}
private void writeFieldExpression(RustFieldExpression fieldExpression) throws IOException {
- this.writeExpression(fieldExpression.containerOperand());
+ this.writeSubExpression(fieldExpression.containerOperand(), precedence(fieldExpression), true);
this.writer.write(".");
this.writeIdentifier(fieldExpression.fieldName());
}
private void writeIfExpression(RustIfExpression ifExpression) throws IOException {
this.writer.write("if ");
- this.writeExpression(ifExpression.condition());
+ this.writeTopLevelExpression(ifExpression.condition());
this.writer.write(" ");
this.writeBlockExpression(ifExpression.ifTrue());
this.writer.write(" else ");
@@ -453,7 +518,7 @@ public class RustWriter implements AutoCloseable {
case BORROW_MUTABLE -> "&mut ";
case DEREFERENCE -> "*";
});
- this.writeExpression(prefixExpression.operand());
+ this.writeSubExpression(prefixExpression.operand(), precedence(prefixExpression), true);
}
private void writeStructExpression(RustStructExpression structExpression) throws IOException {
@@ -467,7 +532,7 @@ public class RustWriter implements AutoCloseable {
fieldValue -> {
this.writeIdentifier(fieldValue.name());
this.writer.write(": ");
- this.writeExpression(fieldValue.value());
+ this.writeTopLevelExpression(fieldValue.value());
},
() -> {
this.writer.write(", ");
@@ -482,7 +547,11 @@ public class RustWriter implements AutoCloseable {
private void writeTryPropagationExpression(
RustTryPropagationExpression tryPropagationExpression
) throws IOException {
- this.writeExpression(tryPropagationExpression.operand());
+ this.writeSubExpression(
+ tryPropagationExpression.operand(),
+ precedence(tryPropagationExpression),
+ true
+ );
this.writer.write("?");
}
@@ -491,12 +560,12 @@ public class RustWriter implements AutoCloseable {
) throws IOException {
this.writer.write("(");
if (tupleExpression.elements().size() == 1) {
- this.writeExpression(tupleExpression.elements().getFirst());
+ this.writeTopLevelExpression(tupleExpression.elements().getFirst());
this.writer.write(",");
} else if (tupleExpression.elements().size() > 1) {
writeWithSeparator(
tupleExpression.elements(),
- this::writeExpression,
+ this::writeTopLevelExpression,
() -> {
this.writer.write(", ");
}
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 862d08b..4279f49 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
@@ -403,7 +403,7 @@ public class RustWriterTests {
.withLengthOperand(RustIntegerLiteral.arbitrary().withValue(47))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("[42; 47]"));
}
@@ -418,7 +418,7 @@ public class RustWriterTests {
.withRight(RustPath.of("y"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("x + y"));
}
@@ -431,7 +431,7 @@ public class RustWriterTests {
.withRight(RustPath.of("y"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("x == y"));
}
@@ -444,7 +444,7 @@ public class RustWriterTests {
.withRight(RustPath.of("y"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("x != y"));
}
@@ -455,7 +455,7 @@ public class RustWriterTests {
public void blockExpressionWithNoStatementsAndNoFinalOperand() throws IOException {
var rust = RustBlockExpression.arbitrary().build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("""
{
@@ -473,7 +473,7 @@ public class RustWriterTests {
)
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("""
{
@@ -488,7 +488,7 @@ public class RustWriterTests {
.withFinalOperand(RustBoolLiteral.arbitrary().withValue(true))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("""
{
@@ -508,7 +508,7 @@ public class RustWriterTests {
.withFinalOperand(RustBoolLiteral.arbitrary().withValue(true))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("""
{
@@ -524,7 +524,7 @@ public class RustWriterTests {
public void trueLiteral() throws IOException {
var rust = RustBoolLiteral.arbitrary().withValue(true).build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("true"));
}
@@ -533,7 +533,7 @@ public class RustWriterTests {
public void falseLiteral() throws IOException {
var rust = RustBoolLiteral.arbitrary().withValue(false).build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("false"));
}
@@ -546,7 +546,7 @@ public class RustWriterTests {
.withFunction(RustPath.of("f"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("f()"));
}
@@ -559,7 +559,7 @@ public class RustWriterTests {
.addArg(RustPath.of("y"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("f(x, y)"));
}
@@ -573,7 +573,7 @@ public class RustWriterTests {
.withFieldName(RustIdentifier.of("y"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("x.y"));
}
@@ -588,7 +588,7 @@ public class RustWriterTests {
.withIfFalse(RustBlockExpression.arbitrary().withFinalOperand(RustPath.of("c")))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("""
if a {
@@ -604,7 +604,7 @@ public class RustWriterTests {
public void integerLiteralWithoutExplicitType() throws IOException {
var rust = RustIntegerLiteral.arbitrary().withValue(123).build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("123"));
}
@@ -616,7 +616,7 @@ public class RustWriterTests {
.withType(RustIdentifier.of("i32"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("123i32"));
}
@@ -630,7 +630,7 @@ public class RustWriterTests {
.withOperand(RustPath.of("x"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("&x"));
}
@@ -642,7 +642,7 @@ public class RustWriterTests {
.withOperand(RustPath.of("x"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("&mut x"));
}
@@ -654,7 +654,7 @@ public class RustWriterTests {
.withOperand(RustPath.of("x"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("*x"));
}
@@ -667,7 +667,7 @@ public class RustWriterTests {
.withStructPath(RustPath.of("A"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("A {}"));
}
@@ -688,7 +688,7 @@ public class RustWriterTests {
)
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("X { a: b, c: d }"));
}
@@ -701,7 +701,7 @@ public class RustWriterTests {
.withOperand(RustPath.of("a"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("a?"));
}
@@ -712,7 +712,7 @@ public class RustWriterTests {
public void tupleExpressionUnit() throws IOException {
var rust = RustTupleExpression.arbitrary().build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("()"));
}
@@ -723,7 +723,7 @@ public class RustWriterTests {
.addElement(RustPath.of("a"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("(a,)"));
}
@@ -736,11 +736,54 @@ public class RustWriterTests {
.addElement(RustPath.of("c"))
.build();
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("(a, b, c)"));
}
+ // === Precedence handling ===
+
+ @Test
+ public void whenSubExpressionHasLowerPrecedenceThenSubExpressionIsParenthesized() throws IOException {
+ var rust = RustBinaryExpression.arbitrary()
+ .withOperator(RustBinaryOperator.ADD)
+ .withLeft(
+ RustBinaryExpression.arbitrary()
+ .withOperator(RustBinaryOperator.EQUAL)
+ .withLeft(RustPath.of("x"))
+ .withRight(RustPath.of("y"))
+ )
+ .withRight(RustPath.of("z"))
+ .build();
+
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
+
+ assertThat(string, equalTo("(x == y) + z"));
+ }
+
+ @Test
+ public void whenSubExpressionsOfLeftAssociativeBinaryOperationHaveSamePrecedenceThenRightExpressionIsParenthesized() throws IOException {
+ var rust = RustBinaryExpression.arbitrary()
+ .withOperator(RustBinaryOperator.ADD)
+ .withLeft(
+ RustBinaryExpression.arbitrary()
+ .withOperator(RustBinaryOperator.ADD)
+ .withLeft(RustPath.of("a"))
+ .withRight(RustPath.of("b"))
+ )
+ .withRight(
+ RustBinaryExpression.arbitrary()
+ .withOperator(RustBinaryOperator.ADD)
+ .withLeft(RustPath.of("c"))
+ .withRight(RustPath.of("d"))
+ )
+ .build();
+
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
+
+ assertThat(string, equalTo("a + b + (c + d)"));
+ }
+
// == Types ==
@Test
@@ -815,7 +858,7 @@ public class RustWriterTests {
public void pathSegmentsAreSeparatedByDoubleColons() throws IOException {
var rust = RustPath.of("std", "string", "String");
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("std::string::String"));
}
@@ -825,7 +868,7 @@ public class RustWriterTests {
var rust = RustPath.of("std", "collections", "HashMap")
.withArgs(List.of(RustPath.of("i32"), RustPath.of("f64")));
- var string = write(writer -> writer.writeExpression(rust));
+ var string = write(writer -> writer.writeTopLevelExpression(rust));
assertThat(string, equalTo("std::collections::HashMap::<i32, f64>"));
}