Attempt two-pass parsing in SLL prediction and fall back to LL prediction.

Due to grammar ambiguities, we fall back to LL prediction considering contextual ambiguity resolution.

Closes #3757
This commit is contained in:
Mark Paluch
2025-01-30 16:50:58 +01:00
parent 76a0f5e409
commit 6038bbad0c
9 changed files with 77 additions and 77 deletions

View File

@@ -20,6 +20,7 @@ import java.util.Set;
import java.util.function.BiFunction;
import java.util.function.Function;
import org.antlr.v4.runtime.BailErrorStrategy;
import org.antlr.v4.runtime.CharStream;
import org.antlr.v4.runtime.CharStreams;
import org.antlr.v4.runtime.CommonTokenStream;
@@ -28,7 +29,9 @@ import org.antlr.v4.runtime.Parser;
import org.antlr.v4.runtime.ParserRuleContext;
import org.antlr.v4.runtime.TokenStream;
import org.antlr.v4.runtime.atn.PredictionMode;
import org.antlr.v4.runtime.misc.ParseCancellationException;
import org.antlr.v4.runtime.tree.ParseTreeVisitor;
import org.springframework.data.domain.Sort;
import org.springframework.lang.Nullable;
import org.springframework.util.Assert;
@@ -65,9 +68,35 @@ class JpaQueryEnhancer implements QueryEnhancer {
this.projection = tokens.isEmpty() ? "" : new QueryRenderer.TokenRenderer(tokens).render();
}
/**
* Parse the query and return the parser context (AST). This method attempts parsing the query using
* {@link PredictionMode#SLL} first to attempt a fast-path parse without using the context. If that fails, it retries
* using {@link PredictionMode#LL} which is much slower, however it allows for contextual ambiguity resolution.
*/
static <P extends Parser> ParserRuleContext parse(String query, Function<CharStream, Lexer> lexerFactoryFunction,
Function<TokenStream, P> parserFactoryFunction, Function<P, ParserRuleContext> parseFunction) {
P parser = getParser(query, lexerFactoryFunction, parserFactoryFunction);
parser.getInterpreter().setPredictionMode(PredictionMode.SLL);
parser.setErrorHandler(new BailErrorStrategy());
try {
return parseFunction.apply(parser);
} catch (BadJpqlGrammarException | ParseCancellationException e) {
parser = getParser(query, lexerFactoryFunction, parserFactoryFunction);
// fall back to LL(*)-based parsing
parser.getInterpreter().setPredictionMode(PredictionMode.LL);
return parseFunction.apply(parser);
}
}
private static <P extends Parser> P getParser(String query, Function<CharStream, Lexer> lexerFactoryFunction,
Function<TokenStream, P> parserFactoryFunction) {
Lexer lexer = lexerFactoryFunction.apply(CharStreams.fromString(query));
P parser = parserFactoryFunction.apply(new CommonTokenStream(lexer));
@@ -79,11 +108,11 @@ class JpaQueryEnhancer implements QueryEnhancer {
configureParser(query, grammar.toUpperCase(), lexer, parser);
return parseFunction.apply(parser);
return parser;
}
/**
* Apply common configuration (SLL prediction for performance, our own error listeners).
* Apply common configuration.
*
* @param query
* @param lexer
@@ -96,8 +125,6 @@ class JpaQueryEnhancer implements QueryEnhancer {
lexer.removeErrorListeners();
lexer.addErrorListener(errorListener);
parser.getInterpreter().setPredictionMode(PredictionMode.SLL);
parser.removeErrorListeners();
parser.addErrorListener(errorListener);
}
@@ -142,6 +169,13 @@ class JpaQueryEnhancer implements QueryEnhancer {
return EqlQueryParser.parseQuery(query.getQueryString());
}
/**
* @return the parser context (AST) representing the parsed query.
*/
ParserRuleContext getContext() {
return context;
}
/**
* Checks if the select clause has a new constructor instantiation in the JPA query.
*

View File

@@ -17,9 +17,8 @@ package org.springframework.data.jpa.repository.query;
import static org.assertj.core.api.Assertions.*;
import org.antlr.v4.runtime.CharStreams;
import org.antlr.v4.runtime.CommonTokenStream;
import org.junit.jupiter.api.Test;
import org.springframework.data.jpa.repository.query.QueryRenderer.TokenRenderer;
/**
@@ -40,14 +39,9 @@ class EqlComplianceTests {
*/
private static String parseWithoutChanges(String query) {
EqlLexer lexer = new EqlLexer(CharStreams.fromString(query));
EqlParser parser = new EqlParser(new CommonTokenStream(lexer));
JpaQueryEnhancer.EqlQueryParser parser = JpaQueryEnhancer.EqlQueryParser.parseQuery(query);
parser.addErrorListener(new BadJpqlGrammarErrorListener(query));
EqlParser.StartContext parsedQuery = parser.start();
return TokenRenderer.render(new EqlQueryRenderer().visit(parsedQuery));
return TokenRenderer.render(new EqlQueryRenderer().visit(parser.getContext()));
}
private void assertQuery(String query) {

View File

@@ -19,14 +19,13 @@ import static org.assertj.core.api.Assertions.*;
import java.util.stream.Stream;
import org.antlr.v4.runtime.CharStreams;
import org.antlr.v4.runtime.CommonTokenStream;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.Arguments;
import org.junit.jupiter.params.provider.MethodSource;
import org.junit.jupiter.params.provider.ValueSource;
import org.springframework.data.jpa.repository.query.QueryRenderer.TokenRenderer;
/**
@@ -47,14 +46,9 @@ class EqlQueryRendererTests {
*/
private static String parseWithoutChanges(String query) {
EqlLexer lexer = new EqlLexer(CharStreams.fromString(query));
EqlParser parser = new EqlParser(new CommonTokenStream(lexer));
JpaQueryEnhancer.EqlQueryParser parser = JpaQueryEnhancer.EqlQueryParser.parseQuery(query);
parser.addErrorListener(new BadJpqlGrammarErrorListener(query));
EqlParser.StartContext parsedQuery = parser.start();
return TokenRenderer.render(new EqlQueryRenderer().visit(parsedQuery));
return TokenRenderer.render(new EqlQueryRenderer().visit(parser.getContext()));
}
static Stream<Arguments> reservedWords() {

View File

@@ -17,10 +17,9 @@ package org.springframework.data.jpa.repository.query;
import static org.assertj.core.api.Assertions.*;
import org.antlr.v4.runtime.CharStreams;
import org.antlr.v4.runtime.CommonTokenStream;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.springframework.data.jpa.repository.query.QueryRenderer.TokenRenderer;
/**
@@ -37,14 +36,9 @@ class EqlSpecificationTests {
private static String parseWithoutChanges(String query) {
EqlLexer lexer = new EqlLexer(CharStreams.fromString(query));
EqlParser parser = new EqlParser(new CommonTokenStream(lexer));
JpaQueryEnhancer.EqlQueryParser parser = JpaQueryEnhancer.EqlQueryParser.parseQuery(query);
parser.addErrorListener(new BadJpqlGrammarErrorListener(query));
EqlParser.StartContext parsedQuery = parser.start();
return TokenRenderer.render(new EqlQueryRenderer().visit(parsedQuery));
return TokenRenderer.render(new EqlQueryRenderer().visit(parser.getContext()));
}
private void assertQuery(String query) {

View File

@@ -19,8 +19,6 @@ import static org.assertj.core.api.Assertions.*;
import java.util.stream.Stream;
import org.antlr.v4.runtime.CharStreams;
import org.antlr.v4.runtime.CommonTokenStream;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.params.ParameterizedTest;
@@ -50,14 +48,9 @@ class HqlQueryRendererTests {
*/
private static String parseWithoutChanges(String query) {
HqlLexer lexer = new HqlLexer(CharStreams.fromString(query));
HqlParser parser = new HqlParser(new CommonTokenStream(lexer));
JpaQueryEnhancer.HqlQueryParser parser = JpaQueryEnhancer.HqlQueryParser.parseQuery(query);
parser.addErrorListener(new BadJpqlGrammarErrorListener(query));
HqlParser.StartContext parsedQuery = parser.start();
QueryTokenStream tokens = new HqlQueryRenderer().visit(parsedQuery);
QueryTokenStream tokens = new HqlQueryRenderer().visit(parser.getContext());
return QueryRenderer.from(tokens).render();
}
@@ -1891,6 +1884,22 @@ class HqlQueryRendererTests {
""");
}
@Test // GH-3757
void arithmeticDate() {
assertQuery("SELECT a FROM foo a WHERE (cast(a.createdAt as date) - CURRENT_DATE()) BY day - 2 = 0");
assertQuery("SELECT a FROM foo a WHERE (cast(a.createdAt as date) - CURRENT_DATE()) BY day - 2 = 0");
assertQuery("SELECT a FROM foo a WHERE (cast(a.createdAt as date)) BY day - 2 = 0");
assertQuery("SELECT f.start - 1 minute FROM foo f");
assertQuery("SELECT f FROM foo f WHERE (cast(f.start as date) - CURRENT_DATE()) BY day - 2 = 0");
assertQuery("SELECT 1 week - 1 day FROM foo f");
assertQuery("SELECT f.birthday - local date day FROM foo f");
assertQuery("SELECT local datetime - f.birthday FROM foo f");
assertQuery("SELECT (1 year) by day FROM foo f");
}
@ParameterizedTest // GH-3342
@ValueSource(
strings = { "select 1 from User", "select -1 from User", "select +1 from User", "select +1 * -100 from User",

View File

@@ -17,12 +17,11 @@ package org.springframework.data.jpa.repository.query;
import static org.assertj.core.api.Assertions.*;
import org.antlr.v4.runtime.CharStreams;
import org.antlr.v4.runtime.CommonTokenStream;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.ValueSource;
import org.springframework.data.jpa.repository.query.QueryRenderer.TokenRenderer;
/**
@@ -43,14 +42,9 @@ class HqlSpecificationTests {
private static String parseWithoutChanges(String query) {
HqlLexer lexer = new HqlLexer(CharStreams.fromString(query));
HqlParser parser = new HqlParser(new CommonTokenStream(lexer));
JpaQueryEnhancer.HqlQueryParser parser = JpaQueryEnhancer.HqlQueryParser.parseQuery(query);
parser.addErrorListener(new BadJpqlGrammarErrorListener(query));
HqlParser.StartContext parsedQuery = parser.start();
return TokenRenderer.render(new HqlQueryRenderer().visit(parsedQuery));
return TokenRenderer.render(new HqlQueryRenderer().visit(parser.getContext()));
}
private void assertQuery(String query) {
@@ -490,7 +484,7 @@ class HqlSpecificationTests {
"from Call c ");
assertQuery("select POSITION(c.number IN 'foo') + 1 AS pos " + //
"from Call c ");
"from Call c ");
}
@Test // GH-3689

View File

@@ -17,8 +17,6 @@ package org.springframework.data.jpa.repository.query;
import static org.assertj.core.api.Assertions.*;
import org.antlr.v4.runtime.CharStreams;
import org.antlr.v4.runtime.CommonTokenStream;
import org.junit.jupiter.api.Test;
/**
@@ -32,14 +30,9 @@ class JpqlComplianceTests {
private static String parseWithoutChanges(String query) {
JpqlLexer lexer = new JpqlLexer(CharStreams.fromString(query));
JpqlParser parser = new JpqlParser(new CommonTokenStream(lexer));
JpaQueryEnhancer.JpqlQueryParser parser = JpaQueryEnhancer.JpqlQueryParser.parseQuery(query);
parser.addErrorListener(new BadJpqlGrammarErrorListener(query));
JpqlParser.StartContext parsedQuery = parser.start();
return QueryRenderer.render(new JpqlQueryRenderer().visit(parsedQuery));
return QueryRenderer.render(new JpqlQueryRenderer().visit(parser.getContext()));
}
private void assertQuery(String query) {

View File

@@ -19,14 +19,13 @@ import static org.assertj.core.api.Assertions.*;
import java.util.stream.Stream;
import org.antlr.v4.runtime.CharStreams;
import org.antlr.v4.runtime.CommonTokenStream;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.Arguments;
import org.junit.jupiter.params.provider.MethodSource;
import org.junit.jupiter.params.provider.ValueSource;
import org.springframework.data.jpa.repository.query.QueryRenderer.TokenRenderer;
/**
@@ -48,14 +47,9 @@ class JpqlQueryRendererTests {
*/
private static String parseWithoutChanges(String query) {
JpqlLexer lexer = new JpqlLexer(CharStreams.fromString(query));
JpqlParser parser = new JpqlParser(new CommonTokenStream(lexer));
JpaQueryEnhancer.JpqlQueryParser parser = JpaQueryEnhancer.JpqlQueryParser.parseQuery(query);
parser.addErrorListener(new BadJpqlGrammarErrorListener(query));
JpqlParser.StartContext parsedQuery = parser.start();
return TokenRenderer.render(new JpqlQueryRenderer().visit(parsedQuery));
return TokenRenderer.render(new JpqlQueryRenderer().visit(parser.getContext()));
}
static Stream<Arguments> reservedWords() {

View File

@@ -17,10 +17,9 @@ package org.springframework.data.jpa.repository.query;
import static org.assertj.core.api.Assertions.*;
import org.antlr.v4.runtime.CharStreams;
import org.antlr.v4.runtime.CommonTokenStream;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.springframework.data.jpa.repository.query.QueryRenderer.TokenRenderer;
/**
@@ -41,14 +40,9 @@ class JpqlSpecificationTests {
*/
private static String parseWithoutChanges(String query) {
JpqlLexer lexer = new JpqlLexer(CharStreams.fromString(query));
JpqlParser parser = new JpqlParser(new CommonTokenStream(lexer));
JpaQueryEnhancer.JpqlQueryParser parser = JpaQueryEnhancer.JpqlQueryParser.parseQuery(query);
parser.addErrorListener(new BadJpqlGrammarErrorListener(query));
JpqlParser.StartContext parsedQuery = parser.start();
return TokenRenderer.render(new JpqlQueryRenderer().visit(parsedQuery));
return TokenRenderer.render(new JpqlQueryRenderer().visit(parser.getContext()));
}
private void assertQuery(String query) {