Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 23 additions & 3 deletions src/main/java/net/sf/jsqlparser/expression/LambdaExpression.java
Original file line number Diff line number Diff line change
Expand Up @@ -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<String> identifiers;
private Expression expression;
private boolean parenthesized;

public LambdaExpression(String identifier, Expression expression) {
this.identifiers = Collections.singletonList(identifier);
Expand All @@ -36,7 +39,18 @@ public static LambdaExpression from(ExpressionList<? extends Expression> 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<String> getIdentifiers() {
Expand All @@ -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<Expression> expressionPrinter) {
if (identifiers.size() == 1 && !parenthesized) {
builder.append(identifiers.get(0));
} else {
int i = 0;
Expand All @@ -68,7 +86,9 @@ public StringBuilder appendTo(StringBuilder builder) {
}
builder.append(" )");
}
return builder.append(" -> ").append(expression);
builder.append(" -> ");
expressionPrinter.accept(expression);
return builder;
}

@Override
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1857,20 +1857,7 @@ public <S> StringBuilder visit(StructType structType, S context) {

@Override
public <S> 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
Expand Down
12 changes: 9 additions & 3 deletions src/main/jjtree/net/sf/jsqlparser/parser/JSqlParserCC.jjt
Original file line number Diff line number Diff line change
Expand Up @@ -1677,6 +1677,12 @@ public class CCJSqlParser extends AbstractJSqlParser<CCJSqlParser> {
}
}

// 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);
Expand Down Expand Up @@ -10390,7 +10396,7 @@ ExpressionList SimpleExpressionList():
(
LOOKAHEAD(2, {!interrupted} ) ","
(
LOOKAHEAD( RelObjectName() "->" ) expr=LambdaExpression()
LOOKAHEAD( RelObjectName() "->", { !isJsonArrowDialect() } ) expr=LambdaExpression()
|
expr=SimpleExpression()
)
Expand Down Expand Up @@ -10453,7 +10459,7 @@ ExpressionList ComplexExpressionList():
|
LOOKAHEAD(2) expr=PostgresNamedFunctionParameter()
|
LOOKAHEAD( RelObjectName() "->" ) expr=LambdaExpression()
LOOKAHEAD( RelObjectName() "->", { !isJsonArrowDialect() } ) expr=LambdaExpression()
|
expr=Expression()
) { expressions.add(expr); }
Expand Down Expand Up @@ -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);
Expand Down
136 changes: 136 additions & 0 deletions src/test/java/net/sf/jsqlparser/expression/LambdaRoundTripTest.java
Original file line number Diff line number Diff line change
@@ -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<LambdaExpression> 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<LambdaExpression> 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 <S> 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<LambdaExpression> lambdas(PlainSelect select) {
List<LambdaExpression> result = new ArrayList<>();
select.getSelectItem(0).getExpression().accept(new ExpressionVisitorAdapter<Void>() {
@Override
public <S> Void visit(LambdaExpression expression, S context) {
result.add(expression);
return super.visit(expression, context);
}
});
return result;
}
}
Original file line number Diff line number Diff line change
@@ -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
Loading