From ab1100148db089f70bfbc7fd0459429d9453a945 Mon Sep 17 00:00:00 2001 From: minleejae Date: Fri, 18 Sep 2026 13:07:14 +0900 Subject: [PATCH] fix: preserve lambda grouping and dialect JSON arrows --- .../expression/LambdaExpression.java | 26 +++- .../util/deparser/ExpressionDeParser.java | 15 +- .../net/sf/jsqlparser/parser/JSqlParserCC.jjt | 12 +- .../expression/LambdaRoundTripTest.java | 136 ++++++++++++++++++ .../expression/arrow-dialect-cases.tsv | 13 ++ 5 files changed, 182 insertions(+), 20 deletions(-) create mode 100644 src/test/java/net/sf/jsqlparser/expression/LambdaRoundTripTest.java create mode 100644 src/test/resources/net/sf/jsqlparser/expression/arrow-dialect-cases.tsv diff --git a/src/main/java/net/sf/jsqlparser/expression/LambdaExpression.java b/src/main/java/net/sf/jsqlparser/expression/LambdaExpression.java index e2819060f5..397a4a4320 100644 --- a/src/main/java/net/sf/jsqlparser/expression/LambdaExpression.java +++ b/src/main/java/net/sf/jsqlparser/expression/LambdaExpression.java @@ -10,15 +10,18 @@ package net.sf.jsqlparser.expression; import net.sf.jsqlparser.expression.operators.relational.ExpressionList; +import net.sf.jsqlparser.expression.operators.relational.ParenthesedExpressionList; import net.sf.jsqlparser.parser.ASTNodeAccessImpl; import java.util.ArrayList; import java.util.Collections; import java.util.List; +import java.util.function.Consumer; public class LambdaExpression extends ASTNodeAccessImpl implements Expression { private List identifiers; private Expression expression; + private boolean parenthesized; public LambdaExpression(String identifier, Expression expression) { this.identifiers = Collections.singletonList(identifier); @@ -36,7 +39,18 @@ public static LambdaExpression from(ExpressionList express for (Expression variable : expressionList) { identifiers.add(variable.toString()); } - return new LambdaExpression(identifiers, expression); + return new LambdaExpression(identifiers, expression) + .setParenthesized(expressionList instanceof ParenthesedExpressionList); + } + + /** Whether the parameter list was explicitly parenthesized, including a single parameter. */ + public boolean isParenthesized() { + return parenthesized; + } + + public LambdaExpression setParenthesized(boolean parenthesized) { + this.parenthesized = parenthesized; + return this; } public List getIdentifiers() { @@ -58,7 +72,11 @@ public LambdaExpression setExpression(Expression expression) { } public StringBuilder appendTo(StringBuilder builder) { - if (identifiers.size() == 1) { + return appendTo(builder, builder::append); + } + + public StringBuilder appendTo(StringBuilder builder, Consumer expressionPrinter) { + if (identifiers.size() == 1 && !parenthesized) { builder.append(identifiers.get(0)); } else { int i = 0; @@ -68,7 +86,9 @@ public StringBuilder appendTo(StringBuilder builder) { } builder.append(" )"); } - return builder.append(" -> ").append(expression); + builder.append(" -> "); + expressionPrinter.accept(expression); + return builder; } @Override diff --git a/src/main/java/net/sf/jsqlparser/util/deparser/ExpressionDeParser.java b/src/main/java/net/sf/jsqlparser/util/deparser/ExpressionDeParser.java index 6ffd4cbc29..c076dd6953 100644 --- a/src/main/java/net/sf/jsqlparser/util/deparser/ExpressionDeParser.java +++ b/src/main/java/net/sf/jsqlparser/util/deparser/ExpressionDeParser.java @@ -1857,20 +1857,7 @@ public StringBuilder visit(StructType structType, S context) { @Override public StringBuilder visit(LambdaExpression lambdaExpression, S context) { - if (lambdaExpression.getIdentifiers().size() == 1) { - builder.append(lambdaExpression.getIdentifiers().get(0)); - } else { - int i = 0; - builder.append("( "); - for (String s : lambdaExpression.getIdentifiers()) { - builder.append(i++ > 0 ? ", " : "").append(s); - } - builder.append(" )"); - } - - builder.append(" -> "); - lambdaExpression.getExpression().accept(this, context); - return builder; + return lambdaExpression.appendTo(builder, expression -> expression.accept(this, context)); } @Override diff --git a/src/main/jjtree/net/sf/jsqlparser/parser/JSqlParserCC.jjt b/src/main/jjtree/net/sf/jsqlparser/parser/JSqlParserCC.jjt index fd6b41f64b..0a3291e98b 100644 --- a/src/main/jjtree/net/sf/jsqlparser/parser/JSqlParserCC.jjt +++ b/src/main/jjtree/net/sf/jsqlparser/parser/JSqlParserCC.jjt @@ -1677,6 +1677,12 @@ public class CCJSqlParser extends AbstractJSqlParser { } } + // In these dialects -> always denotes JSON access, even after a comma or parentheses. + private boolean isJsonArrowDialect() { + return Dialect.POSTGRESQL.name().equals(getAsString(Feature.dialect)) + || isMySqlDialect(); + } + private boolean isMySqlDialect() { String dialect = getAsString(Feature.dialect); return Dialect.MYSQL.name().equals(dialect) || Dialect.MARIADB.name().equals(dialect); @@ -10390,7 +10396,7 @@ ExpressionList SimpleExpressionList(): ( LOOKAHEAD(2, {!interrupted} ) "," ( - LOOKAHEAD( RelObjectName() "->" ) expr=LambdaExpression() + LOOKAHEAD( RelObjectName() "->", { !isJsonArrowDialect() } ) expr=LambdaExpression() | expr=SimpleExpression() ) @@ -10453,7 +10459,7 @@ ExpressionList ComplexExpressionList(): | LOOKAHEAD(2) expr=PostgresNamedFunctionParameter() | - LOOKAHEAD( RelObjectName() "->" ) expr=LambdaExpression() + LOOKAHEAD( RelObjectName() "->", { !isJsonArrowDialect() } ) expr=LambdaExpression() | expr=Expression() ) { expressions.add(expr); } @@ -10852,7 +10858,7 @@ Expression PrimaryExpression() #PrimaryExpression: // SELECT map_filter(my_column, (k,v) -> v.my_inner_column = 'some_value') // First-arg form (issue #2195): array_map((x,y,z) -> x + y, ...) ( - LOOKAHEAD(2) "->" + LOOKAHEAD(2, { !isJsonArrowDialect() }) "->" retval = Expression() { retval = LambdaExpression.from(list, retval); diff --git a/src/test/java/net/sf/jsqlparser/expression/LambdaRoundTripTest.java b/src/test/java/net/sf/jsqlparser/expression/LambdaRoundTripTest.java new file mode 100644 index 0000000000..7c5f94dc80 --- /dev/null +++ b/src/test/java/net/sf/jsqlparser/expression/LambdaRoundTripTest.java @@ -0,0 +1,136 @@ +/*- + * #%L + * JSQLParser library + * %% + * Copyright (C) 2004 - 2026 JSQLParser + * %% + * Dual licensed under GNU LGPL 2.1 or Apache License 2.0 + * #L% + */ +package net.sf.jsqlparser.expression; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.util.ArrayList; +import java.util.List; +import net.sf.jsqlparser.JSQLParserException; +import net.sf.jsqlparser.parser.AbstractJSqlParser.Dialect; +import net.sf.jsqlparser.parser.CCJSqlParserUtil; +import net.sf.jsqlparser.statement.select.PlainSelect; +import net.sf.jsqlparser.util.deparser.ExpressionDeParser; +import net.sf.jsqlparser.util.deparser.StatementDeParser; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvFileSource; +import org.junit.jupiter.params.provider.EnumSource; +import org.junit.jupiter.params.provider.ValueSource; + +class LambdaRoundTripTest { + + @ParameterizedTest + @ValueSource(strings = {"SELECT arrayMap((x) -> x * 2, [1, 2])", + "SELECT list_transform([1, 2], (x) -> x + 1)", + "SELECT f(1, [1, 2], (x) -> CASE WHEN x > 1 THEN x ELSE 0 END)", + "SELECT arrayMap((x, y) -> x + y, [1], [2])", + "SELECT list_transform([1], (x) -> list_transform([2], (y) -> x + y))"}) + void preservesParameterGroupingAndNestedAst(String sql) throws JSQLParserException { + PlainSelect select = (PlainSelect) CCJSqlParserUtil.parse(sql); + List before = lambdas(select); + assertFalse(before.isEmpty()); + assertTrue(before.stream().allMatch(LambdaExpression::isParenthesized)); + StringBuilder visitorSql = new StringBuilder(); + select.accept(new StatementDeParser(visitorSql)); + for (String rendered : List.of(select.toString(), visitorSql.toString())) { + PlainSelect reparsed = (PlainSelect) CCJSqlParserUtil.parse(rendered); + List after = lambdas(reparsed); + assertEquals(before.size(), after.size(), rendered); + for (int i = 0; i < before.size(); i++) { + assertEquals(before.get(i).getIdentifiers(), after.get(i).getIdentifiers()); + assertEquals(before.get(i).getExpression().getClass(), + after.get(i).getExpression().getClass()); + assertTrue(after.get(i).isParenthesized()); + } + assertEquals(select.toString(), reparsed.toString()); + } + } + + @Test + void keepsUnparenthesizedSingleParameterRendering() throws JSQLParserException { + PlainSelect select = (PlainSelect) CCJSqlParserUtil.parse( + "SELECT list_transform([1], x -> x + 1)"); + assertFalse(lambdas(select).get(0).isParenthesized()); + assertEquals("SELECT list_transform([1], x -> x + 1)", select.toString()); + LambdaExpression constructed = new LambdaExpression("x", new LongValue(1)); + assertEquals("x -> 1", constructed.toString()); + constructed.setParenthesized(true); + assertEquals("( x ) -> 1", constructed.toString()); + } + + @Test + void sharedRendererStillVisitsAndRewritesLambdaBody() throws JSQLParserException { + PlainSelect select = (PlainSelect) CCJSqlParserUtil.parse( + "SELECT arrayMap((x) -> x + 1, [2])"); + StringBuilder sql = new StringBuilder(); + ExpressionDeParser expressions = new ExpressionDeParser() { + @Override + public StringBuilder visit(LongValue value, S context) { + return getBuilder().append(value.getValue() + 10); + } + }; + select.accept(new StatementDeParser(expressions, + new net.sf.jsqlparser.util.deparser.SelectDeParser(), sql)); + assertEquals("SELECT arrayMap(( x ) -> x + 11, [12])", sql.toString()); + assertEquals("SELECT arrayMap(( x ) -> x + 1, [2])", select.toString()); + assertEquals(1, lambdas((PlainSelect) CCJSqlParserUtil.parse(sql.toString())).size()); + } + + // All fixture statements execute on PostgreSQL 18.6 or MySQL 8.4.11, respectively. + @ParameterizedTest + @CsvFileSource(resources = "/net/sf/jsqlparser/expression/arrow-dialect-cases.tsv", + delimiter = '\t') + void jsonArrowRemainsJsonInEveryArgumentPosition(Dialect dialect, String sql) + throws JSQLParserException { + for (boolean complex : List.of(false, true)) { + PlainSelect select = parseJson(sql, dialect, complex); + StringBuilder visitorSql = new StringBuilder(); + select.accept(new StatementDeParser(visitorSql)); + for (String rendered : List.of(sql, select.toString(), visitorSql.toString())) { + PlainSelect reparsed = parseJson(rendered, dialect, complex); + assertTrue(lambdas(reparsed).isEmpty(), rendered); + Function function = assertInstanceOf(Function.class, + reparsed.getSelectItem(0).getExpression()); + assertTrue(function.getParameters().stream() + .anyMatch(JsonExpression.class::isInstance), rendered); + } + } + } + + @ParameterizedTest + @EnumSource(value = Dialect.class, names = {"MYSQL", "POSTGRESQL"}) + void rejectsMissingJsonOperand(Dialect dialect) { + assertThrows(JSQLParserException.class, + () -> parseJson("SELECT COALESCE(NULL, payload ->)", dialect, true)); + } + + private static PlainSelect parseJson(String sql, Dialect dialect, boolean complex) + throws JSQLParserException { + return (PlainSelect) CCJSqlParserUtil.parse(sql, + parser -> parser.withDialect(dialect).withAllowComplexParsing(complex)); + } + + private static List lambdas(PlainSelect select) { + List result = new ArrayList<>(); + select.getSelectItem(0).getExpression().accept(new ExpressionVisitorAdapter() { + @Override + public Void visit(LambdaExpression expression, S context) { + result.add(expression); + return super.visit(expression, context); + } + }); + return result; + } +} diff --git a/src/test/resources/net/sf/jsqlparser/expression/arrow-dialect-cases.tsv b/src/test/resources/net/sf/jsqlparser/expression/arrow-dialect-cases.tsv new file mode 100644 index 0000000000..4c68de1fbd --- /dev/null +++ b/src/test/resources/net/sf/jsqlparser/expression/arrow-dialect-cases.tsv @@ -0,0 +1,13 @@ +POSTGRESQL SELECT COALESCE(payload -> 'a', payload -> 'b') FROM arrow_inputs +POSTGRESQL SELECT COALESCE((payload) -> key, payload -> 0) FROM arrow_inputs +POSTGRESQL SELECT COALESCE(payload -> -1, (payload) -> (idx + 1)) FROM arrow_inputs +POSTGRESQL SELECT COALESCE(payload -> 'a' -> 'b', payload -> key) FROM arrow_inputs +POSTGRESQL SELECT COALESCE(payload ->> 'a', payload ->> key) FROM arrow_inputs +POSTGRESQL SELECT COALESCE(NULL::jsonb, payload -> key) FROM arrow_inputs +POSTGRESQL SELECT COALESCE(NULL::jsonb, (payload) -> key) FROM arrow_inputs +POSTGRESQL SELECT COALESCE(payload -> 'a', (payload) -> 'a', payload -> key) FROM arrow_inputs +MYSQL SELECT COALESCE(payload -> '$.a', payload -> '$.b') FROM arrow_inputs +MYSQL SELECT COALESCE(NULL, payload -> '$.a') FROM arrow_inputs +MYSQL SELECT COALESCE(payload ->> '$.a', payload ->> '$.b') FROM arrow_inputs +MYSQL SELECT COALESCE(NULL, payload -> '$[0]', payload -> '$[last]') FROM arrow_inputs +MYSQL SELECT COALESCE(i.payload -> '$.a', i.payload -> '$.b') FROM arrow_inputs i