Skip to content

Commit a8ef0b2

Browse files
authored
fix: preserve lambda grouping and dialect JSON arrows (#2646)
1 parent d1dbca1 commit a8ef0b2

5 files changed

Lines changed: 182 additions & 20 deletions

File tree

src/main/java/net/sf/jsqlparser/expression/LambdaExpression.java

Lines changed: 23 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -10,15 +10,18 @@
1010
package net.sf.jsqlparser.expression;
1111

1212
import net.sf.jsqlparser.expression.operators.relational.ExpressionList;
13+
import net.sf.jsqlparser.expression.operators.relational.ParenthesedExpressionList;
1314
import net.sf.jsqlparser.parser.ASTNodeAccessImpl;
1415

1516
import java.util.ArrayList;
1617
import java.util.Collections;
1718
import java.util.List;
19+
import java.util.function.Consumer;
1820

1921
public class LambdaExpression extends ASTNodeAccessImpl implements Expression {
2022
private List<String> identifiers;
2123
private Expression expression;
24+
private boolean parenthesized;
2225

2326
public LambdaExpression(String identifier, Expression expression) {
2427
this.identifiers = Collections.singletonList(identifier);
@@ -36,7 +39,18 @@ public static LambdaExpression from(ExpressionList<? extends Expression> express
3639
for (Expression variable : expressionList) {
3740
identifiers.add(variable.toString());
3841
}
39-
return new LambdaExpression(identifiers, expression);
42+
return new LambdaExpression(identifiers, expression)
43+
.setParenthesized(expressionList instanceof ParenthesedExpressionList);
44+
}
45+
46+
/** Whether the parameter list was explicitly parenthesized, including a single parameter. */
47+
public boolean isParenthesized() {
48+
return parenthesized;
49+
}
50+
51+
public LambdaExpression setParenthesized(boolean parenthesized) {
52+
this.parenthesized = parenthesized;
53+
return this;
4054
}
4155

4256
public List<String> getIdentifiers() {
@@ -58,7 +72,11 @@ public LambdaExpression setExpression(Expression expression) {
5872
}
5973

6074
public StringBuilder appendTo(StringBuilder builder) {
61-
if (identifiers.size() == 1) {
75+
return appendTo(builder, builder::append);
76+
}
77+
78+
public StringBuilder appendTo(StringBuilder builder, Consumer<Expression> expressionPrinter) {
79+
if (identifiers.size() == 1 && !parenthesized) {
6280
builder.append(identifiers.get(0));
6381
} else {
6482
int i = 0;
@@ -68,7 +86,9 @@ public StringBuilder appendTo(StringBuilder builder) {
6886
}
6987
builder.append(" )");
7088
}
71-
return builder.append(" -> ").append(expression);
89+
builder.append(" -> ");
90+
expressionPrinter.accept(expression);
91+
return builder;
7292
}
7393

7494
@Override

src/main/java/net/sf/jsqlparser/util/deparser/ExpressionDeParser.java

Lines changed: 1 addition & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -1857,20 +1857,7 @@ public <S> StringBuilder visit(StructType structType, S context) {
18571857

18581858
@Override
18591859
public <S> StringBuilder visit(LambdaExpression lambdaExpression, S context) {
1860-
if (lambdaExpression.getIdentifiers().size() == 1) {
1861-
builder.append(lambdaExpression.getIdentifiers().get(0));
1862-
} else {
1863-
int i = 0;
1864-
builder.append("( ");
1865-
for (String s : lambdaExpression.getIdentifiers()) {
1866-
builder.append(i++ > 0 ? ", " : "").append(s);
1867-
}
1868-
builder.append(" )");
1869-
}
1870-
1871-
builder.append(" -> ");
1872-
lambdaExpression.getExpression().accept(this, context);
1873-
return builder;
1860+
return lambdaExpression.appendTo(builder, expression -> expression.accept(this, context));
18741861
}
18751862

18761863
@Override

src/main/jjtree/net/sf/jsqlparser/parser/JSqlParserCC.jjt

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1677,6 +1677,12 @@ public class CCJSqlParser extends AbstractJSqlParser<CCJSqlParser> {
16771677
}
16781678
}
16791679

1680+
// In these dialects -> always denotes JSON access, even after a comma or parentheses.
1681+
private boolean isJsonArrowDialect() {
1682+
return Dialect.POSTGRESQL.name().equals(getAsString(Feature.dialect))
1683+
|| isMySqlDialect();
1684+
}
1685+
16801686
private boolean isMySqlDialect() {
16811687
String dialect = getAsString(Feature.dialect);
16821688
return Dialect.MYSQL.name().equals(dialect) || Dialect.MARIADB.name().equals(dialect);
@@ -10420,7 +10426,7 @@ ExpressionList SimpleExpressionList():
1042010426
(
1042110427
LOOKAHEAD(2, {!interrupted} ) ","
1042210428
(
10423-
LOOKAHEAD( RelObjectName() "->" ) expr=LambdaExpression()
10429+
LOOKAHEAD( RelObjectName() "->", { !isJsonArrowDialect() } ) expr=LambdaExpression()
1042410430
|
1042510431
expr=SimpleExpression()
1042610432
)
@@ -10483,7 +10489,7 @@ ExpressionList ComplexExpressionList():
1048310489
|
1048410490
LOOKAHEAD(2) expr=PostgresNamedFunctionParameter()
1048510491
|
10486-
LOOKAHEAD( RelObjectName() "->" ) expr=LambdaExpression()
10492+
LOOKAHEAD( RelObjectName() "->", { !isJsonArrowDialect() } ) expr=LambdaExpression()
1048710493
|
1048810494
expr=Expression()
1048910495
) { expressions.add(expr); }
@@ -10888,7 +10894,7 @@ Expression PrimaryExpression() #PrimaryExpression:
1088810894
// SELECT map_filter(my_column, (k,v) -> v.my_inner_column = 'some_value')
1088910895
// First-arg form (issue #2195): array_map((x,y,z) -> x + y, ...)
1089010896
(
10891-
LOOKAHEAD(2) "->"
10897+
LOOKAHEAD(2, { !isJsonArrowDialect() }) "->"
1089210898
retval = Expression()
1089310899
{
1089410900
retval = LambdaExpression.from(list, retval);
Lines changed: 136 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,136 @@
1+
/*-
2+
* #%L
3+
* JSQLParser library
4+
* %%
5+
* Copyright (C) 2004 - 2026 JSQLParser
6+
* %%
7+
* Dual licensed under GNU LGPL 2.1 or Apache License 2.0
8+
* #L%
9+
*/
10+
package net.sf.jsqlparser.expression;
11+
12+
import static org.junit.jupiter.api.Assertions.assertEquals;
13+
import static org.junit.jupiter.api.Assertions.assertFalse;
14+
import static org.junit.jupiter.api.Assertions.assertInstanceOf;
15+
import static org.junit.jupiter.api.Assertions.assertThrows;
16+
import static org.junit.jupiter.api.Assertions.assertTrue;
17+
18+
import java.util.ArrayList;
19+
import java.util.List;
20+
import net.sf.jsqlparser.JSQLParserException;
21+
import net.sf.jsqlparser.parser.AbstractJSqlParser.Dialect;
22+
import net.sf.jsqlparser.parser.CCJSqlParserUtil;
23+
import net.sf.jsqlparser.statement.select.PlainSelect;
24+
import net.sf.jsqlparser.util.deparser.ExpressionDeParser;
25+
import net.sf.jsqlparser.util.deparser.StatementDeParser;
26+
import org.junit.jupiter.api.Test;
27+
import org.junit.jupiter.params.ParameterizedTest;
28+
import org.junit.jupiter.params.provider.CsvFileSource;
29+
import org.junit.jupiter.params.provider.EnumSource;
30+
import org.junit.jupiter.params.provider.ValueSource;
31+
32+
class LambdaRoundTripTest {
33+
34+
@ParameterizedTest
35+
@ValueSource(strings = {"SELECT arrayMap((x) -> x * 2, [1, 2])",
36+
"SELECT list_transform([1, 2], (x) -> x + 1)",
37+
"SELECT f(1, [1, 2], (x) -> CASE WHEN x > 1 THEN x ELSE 0 END)",
38+
"SELECT arrayMap((x, y) -> x + y, [1], [2])",
39+
"SELECT list_transform([1], (x) -> list_transform([2], (y) -> x + y))"})
40+
void preservesParameterGroupingAndNestedAst(String sql) throws JSQLParserException {
41+
PlainSelect select = (PlainSelect) CCJSqlParserUtil.parse(sql);
42+
List<LambdaExpression> before = lambdas(select);
43+
assertFalse(before.isEmpty());
44+
assertTrue(before.stream().allMatch(LambdaExpression::isParenthesized));
45+
StringBuilder visitorSql = new StringBuilder();
46+
select.accept(new StatementDeParser(visitorSql));
47+
for (String rendered : List.of(select.toString(), visitorSql.toString())) {
48+
PlainSelect reparsed = (PlainSelect) CCJSqlParserUtil.parse(rendered);
49+
List<LambdaExpression> after = lambdas(reparsed);
50+
assertEquals(before.size(), after.size(), rendered);
51+
for (int i = 0; i < before.size(); i++) {
52+
assertEquals(before.get(i).getIdentifiers(), after.get(i).getIdentifiers());
53+
assertEquals(before.get(i).getExpression().getClass(),
54+
after.get(i).getExpression().getClass());
55+
assertTrue(after.get(i).isParenthesized());
56+
}
57+
assertEquals(select.toString(), reparsed.toString());
58+
}
59+
}
60+
61+
@Test
62+
void keepsUnparenthesizedSingleParameterRendering() throws JSQLParserException {
63+
PlainSelect select = (PlainSelect) CCJSqlParserUtil.parse(
64+
"SELECT list_transform([1], x -> x + 1)");
65+
assertFalse(lambdas(select).get(0).isParenthesized());
66+
assertEquals("SELECT list_transform([1], x -> x + 1)", select.toString());
67+
LambdaExpression constructed = new LambdaExpression("x", new LongValue(1));
68+
assertEquals("x -> 1", constructed.toString());
69+
constructed.setParenthesized(true);
70+
assertEquals("( x ) -> 1", constructed.toString());
71+
}
72+
73+
@Test
74+
void sharedRendererStillVisitsAndRewritesLambdaBody() throws JSQLParserException {
75+
PlainSelect select = (PlainSelect) CCJSqlParserUtil.parse(
76+
"SELECT arrayMap((x) -> x + 1, [2])");
77+
StringBuilder sql = new StringBuilder();
78+
ExpressionDeParser expressions = new ExpressionDeParser() {
79+
@Override
80+
public <S> StringBuilder visit(LongValue value, S context) {
81+
return getBuilder().append(value.getValue() + 10);
82+
}
83+
};
84+
select.accept(new StatementDeParser(expressions,
85+
new net.sf.jsqlparser.util.deparser.SelectDeParser(), sql));
86+
assertEquals("SELECT arrayMap(( x ) -> x + 11, [12])", sql.toString());
87+
assertEquals("SELECT arrayMap(( x ) -> x + 1, [2])", select.toString());
88+
assertEquals(1, lambdas((PlainSelect) CCJSqlParserUtil.parse(sql.toString())).size());
89+
}
90+
91+
// All fixture statements execute on PostgreSQL 18.6 or MySQL 8.4.11, respectively.
92+
@ParameterizedTest
93+
@CsvFileSource(resources = "/net/sf/jsqlparser/expression/arrow-dialect-cases.tsv",
94+
delimiter = '\t')
95+
void jsonArrowRemainsJsonInEveryArgumentPosition(Dialect dialect, String sql)
96+
throws JSQLParserException {
97+
for (boolean complex : List.of(false, true)) {
98+
PlainSelect select = parseJson(sql, dialect, complex);
99+
StringBuilder visitorSql = new StringBuilder();
100+
select.accept(new StatementDeParser(visitorSql));
101+
for (String rendered : List.of(sql, select.toString(), visitorSql.toString())) {
102+
PlainSelect reparsed = parseJson(rendered, dialect, complex);
103+
assertTrue(lambdas(reparsed).isEmpty(), rendered);
104+
Function function = assertInstanceOf(Function.class,
105+
reparsed.getSelectItem(0).getExpression());
106+
assertTrue(function.getParameters().stream()
107+
.anyMatch(JsonExpression.class::isInstance), rendered);
108+
}
109+
}
110+
}
111+
112+
@ParameterizedTest
113+
@EnumSource(value = Dialect.class, names = {"MYSQL", "POSTGRESQL"})
114+
void rejectsMissingJsonOperand(Dialect dialect) {
115+
assertThrows(JSQLParserException.class,
116+
() -> parseJson("SELECT COALESCE(NULL, payload ->)", dialect, true));
117+
}
118+
119+
private static PlainSelect parseJson(String sql, Dialect dialect, boolean complex)
120+
throws JSQLParserException {
121+
return (PlainSelect) CCJSqlParserUtil.parse(sql,
122+
parser -> parser.withDialect(dialect).withAllowComplexParsing(complex));
123+
}
124+
125+
private static List<LambdaExpression> lambdas(PlainSelect select) {
126+
List<LambdaExpression> result = new ArrayList<>();
127+
select.getSelectItem(0).getExpression().accept(new ExpressionVisitorAdapter<Void>() {
128+
@Override
129+
public <S> Void visit(LambdaExpression expression, S context) {
130+
result.add(expression);
131+
return super.visit(expression, context);
132+
}
133+
});
134+
return result;
135+
}
136+
}
Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,13 @@
1+
POSTGRESQL SELECT COALESCE(payload -> 'a', payload -> 'b') FROM arrow_inputs
2+
POSTGRESQL SELECT COALESCE((payload) -> key, payload -> 0) FROM arrow_inputs
3+
POSTGRESQL SELECT COALESCE(payload -> -1, (payload) -> (idx + 1)) FROM arrow_inputs
4+
POSTGRESQL SELECT COALESCE(payload -> 'a' -> 'b', payload -> key) FROM arrow_inputs
5+
POSTGRESQL SELECT COALESCE(payload ->> 'a', payload ->> key) FROM arrow_inputs
6+
POSTGRESQL SELECT COALESCE(NULL::jsonb, payload -> key) FROM arrow_inputs
7+
POSTGRESQL SELECT COALESCE(NULL::jsonb, (payload) -> key) FROM arrow_inputs
8+
POSTGRESQL SELECT COALESCE(payload -> 'a', (payload) -> 'a', payload -> key) FROM arrow_inputs
9+
MYSQL SELECT COALESCE(payload -> '$.a', payload -> '$.b') FROM arrow_inputs
10+
MYSQL SELECT COALESCE(NULL, payload -> '$.a') FROM arrow_inputs
11+
MYSQL SELECT COALESCE(payload ->> '$.a', payload ->> '$.b') FROM arrow_inputs
12+
MYSQL SELECT COALESCE(NULL, payload -> '$[0]', payload -> '$[last]') FROM arrow_inputs
13+
MYSQL SELECT COALESCE(i.payload -> '$.a', i.payload -> '$.b') FROM arrow_inputs i

0 commit comments

Comments
 (0)