diff --git a/src/main/java/org/springframework/data/r2dbc/repository/query/ExpressionQuery.java b/src/main/java/org/springframework/data/r2dbc/repository/query/ExpressionQuery.java index 4065f3ae..64047631 100644 --- a/src/main/java/org/springframework/data/r2dbc/repository/query/ExpressionQuery.java +++ b/src/main/java/org/springframework/data/r2dbc/repository/query/ExpressionQuery.java @@ -17,10 +17,8 @@ package org.springframework.data.r2dbc.repository.query; import java.util.ArrayList; import java.util.List; -import java.util.regex.Matcher; -import java.util.regex.Pattern; -import org.springframework.lang.Nullable; +import org.springframework.data.repository.query.SpelQueryContext; /** * Query using Spring Expression Language to indicate parameter bindings. Queries using SpEL use {@code :#{…}} to @@ -32,18 +30,14 @@ import org.springframework.lang.Nullable; */ class ExpressionQuery { - private static final char CURLY_BRACE_OPEN = '{'; - private static final char CURLY_BRACE_CLOSE = '}'; - private static final String SYNTHETIC_PARAMETER_TEMPLATE = "__synthetic_%d__"; - private static final Pattern EXPRESSION_BINDING_PATTERN = Pattern.compile("[:]#\\{(.*)}"); - private final String query; private final List parameterBindings; private ExpressionQuery(String query, List parameterBindings) { + this.query = query; this.parameterBindings = parameterBindings; } @@ -57,9 +51,17 @@ class ExpressionQuery { public static ExpressionQuery create(String query) { List parameterBindings = new ArrayList<>(); - String rewritten = transformQueryAndCollectExpressionParametersIntoBindings(query, parameterBindings); - return new ExpressionQuery(rewritten, parameterBindings); + SpelQueryContext queryContext = SpelQueryContext.of((counter, expression) -> { + + String parameterName = String.format(SYNTHETIC_PARAMETER_TEMPLATE, counter); + parameterBindings.add(new ParameterBinding(parameterName, expression)); + return parameterName; + }, String::concat); + + SpelQueryContext.SpelExtractor parsed = queryContext.parse(query); + + return new ExpressionQuery(parsed.getQueryString(), parameterBindings); } public String getQuery() { @@ -70,67 +72,6 @@ class ExpressionQuery { return parameterBindings; } - private static String transformQueryAndCollectExpressionParametersIntoBindings(String input, - List bindings) { - - StringBuilder result = new StringBuilder(); - - int startIndex = 0; - int currentPosition = 0; - int parameterIndex = 0; - - while (currentPosition < input.length()) { - - Matcher matcher = findNextBindingOrExpression(input, currentPosition); - - // no expression parameter found - if (matcher == null) { - break; - } - - int exprStart = matcher.start(); - currentPosition = exprStart; - - // eat parameter expression - int curlyBraceOpenCount = 1; - currentPosition += 3; - - while (curlyBraceOpenCount > 0 && currentPosition < input.length()) { - switch (input.charAt(currentPosition++)) { - case CURLY_BRACE_OPEN: - curlyBraceOpenCount++; - break; - case CURLY_BRACE_CLOSE: - curlyBraceOpenCount--; - break; - default: - } - } - - result.append(input.subSequence(startIndex, exprStart)); - - String parameterName = String.format(SYNTHETIC_PARAMETER_TEMPLATE, parameterIndex++); - result.append(':').append(parameterName); - - bindings.add(ParameterBinding.named(parameterName, matcher.group(1))); - - currentPosition = matcher.end(); - startIndex = currentPosition; - } - - return result.append(input.subSequence(currentPosition, input.length())).toString(); - } - - @Nullable - private static Matcher findNextBindingOrExpression(String input, int position) { - - Matcher matcher = EXPRESSION_BINDING_PATTERN.matcher(input); - if (matcher.find(position)) { - return matcher; - } - - return null; - } @Override public String toString() { @@ -153,10 +94,6 @@ class ExpressionQuery { this.parameterName = parameterName; } - static ParameterBinding named(String name, String expression) { - return new ParameterBinding(name, expression); - } - String getExpression() { return expression; } diff --git a/src/test/java/org/springframework/data/r2dbc/repository/query/StringBasedR2dbcQueryUnitTests.java b/src/test/java/org/springframework/data/r2dbc/repository/query/StringBasedR2dbcQueryUnitTests.java index 7cde0801..091bb5c4 100644 --- a/src/test/java/org/springframework/data/r2dbc/repository/query/StringBasedR2dbcQueryUnitTests.java +++ b/src/test/java/org/springframework/data/r2dbc/repository/query/StringBasedR2dbcQueryUnitTests.java @@ -28,6 +28,7 @@ import org.mockito.Mock; import org.mockito.junit.MockitoJUnitRunner; import org.springframework.data.domain.Sort; +import org.springframework.data.geo.Point; import org.springframework.data.projection.ProjectionFactory; import org.springframework.data.projection.SpelAwareProxyProjectionFactory; import org.springframework.data.r2dbc.convert.MappingR2dbcConverter; @@ -237,6 +238,23 @@ public class StringBasedR2dbcQueryUnitTests { verifyNoMoreInteractions(bindSpec); } + @Test // gh-373 + public void bindsMultipleSpelParametersCorrectly() { + + StringBasedR2dbcQuery query = getQueryMethod("queryWithTwoSpelExpressions", Point.class); + R2dbcParameterAccessor accessor = new R2dbcParameterAccessor(query.getQueryMethod(), new Point(1, 2)); + + BindableQuery stringQuery = query.createQuery(accessor); + + assertThat(stringQuery.get()) + .isEqualTo("INSERT IGNORE INTO table (x, y) VALUES (:__synthetic_0__, :__synthetic_1__)"); + assertThat(stringQuery.bind(bindSpec)).isNotNull(); + + verify(bindSpec).bind("__synthetic_0__", 1d); + verify(bindSpec).bind("__synthetic_1__", 2d); + verifyNoMoreInteractions(bindSpec); + } + private StringBasedR2dbcQuery getQueryMethod(String name, Class... args) { Method method = ReflectionUtils.findMethod(SampleRepository.class, name, args); @@ -282,6 +300,9 @@ public class StringBasedR2dbcQueryUnitTests { @Query("SELECT * FROM person WHERE lastname = :name") Person queryWithUnusedParameter(String name, Sort unused); + + @Query("INSERT IGNORE INTO table (x, y) VALUES (:#{#point.x}, :#{#point.y})") + Person queryWithTwoSpelExpressions(@Param("point") Point point); } static class Person {