From fd66f8af2640e911b874d6c4ddb2846412c52576 Mon Sep 17 00:00:00 2001 From: Mark Paluch Date: Fri, 20 Sep 2019 15:50:36 +0200 Subject: [PATCH] #178 - Consider named and indexed parameters for named parameter processing. Named parameters can be provided by name and by index. Repository query methods bind parameters by name if a named parameter can be found. If parameters are bound by index, then the parameter name is looked up by index (index corresponds with the order of parameter name discovery when parsing the query). and bound to the parameter. --- .../r2dbc/core/DefaultDatabaseClient.java | 21 +++++++- .../DefaultReactiveDataAccessStrategy.java | 24 +++++++-- .../r2dbc/core/NamedParameterExpander.java | 11 ++++ .../core/ReactiveDataAccessStrategy.java | 31 ++++++++--- .../query/StringBasedR2dbcQuery.java | 42 +++++++++++++-- ...bstractDatabaseClientIntegrationTests.java | 2 +- .../core/DefaultDatabaseClientUnitTests.java | 27 ++++++++++ .../query/StringBasedR2dbcQueryUnitTests.java | 53 +++++++++++++++++++ 8 files changed, 194 insertions(+), 17 deletions(-) diff --git a/src/main/java/org/springframework/data/r2dbc/core/DefaultDatabaseClient.java b/src/main/java/org/springframework/data/r2dbc/core/DefaultDatabaseClient.java index c90c8ecd..0e87081c 100644 --- a/src/main/java/org/springframework/data/r2dbc/core/DefaultDatabaseClient.java +++ b/src/main/java/org/springframework/data/r2dbc/core/DefaultDatabaseClient.java @@ -345,7 +345,22 @@ class DefaultDatabaseClient implements DatabaseClient, ConnectionAccessor { if (namedParameters) { - PreparedOperation operation = dataAccessStrategy.processNamedParameters(sql, this.byName); + Map remainderByName = new LinkedHashMap<>(this.byName); + Map remainderByIndex = new LinkedHashMap<>(this.byIndex); + PreparedOperation operation = dataAccessStrategy.processNamedParameters(sql, (index, name) -> { + + if (byName.containsKey(name)) { + remainderByName.remove(name); + return byName.get(name); + } + + if (byIndex.containsKey(index)) { + remainderByIndex.remove(index); + return byIndex.get(index); + } + + return null; + }); String expanded = getRequiredSql(operation); if (logger.isTraceEnabled()) { @@ -356,7 +371,9 @@ class DefaultDatabaseClient implements DatabaseClient, ConnectionAccessor { BindTarget bindTarget = new StatementWrapper(statement); operation.bindTo(bindTarget); - bindByIndex(statement, this.byIndex); + + bindByName(statement, remainderByName); + bindByIndex(statement, remainderByIndex); return statement; } diff --git a/src/main/java/org/springframework/data/r2dbc/core/DefaultReactiveDataAccessStrategy.java b/src/main/java/org/springframework/data/r2dbc/core/DefaultReactiveDataAccessStrategy.java index ef2e59c0..5d18989f 100644 --- a/src/main/java/org/springframework/data/r2dbc/core/DefaultReactiveDataAccessStrategy.java +++ b/src/main/java/org/springframework/data/r2dbc/core/DefaultReactiveDataAccessStrategy.java @@ -21,10 +21,12 @@ import io.r2dbc.spi.RowMetadata; import java.util.ArrayList; import java.util.Collection; import java.util.Collections; +import java.util.LinkedHashMap; import java.util.List; import java.util.Map; import java.util.function.BiFunction; +import org.springframework.dao.InvalidDataAccessApiUsageException; import org.springframework.dao.InvalidDataAccessResourceUsageException; import org.springframework.data.convert.CustomConversions.StoreConversions; import org.springframework.data.mapping.context.MappingContext; @@ -254,11 +256,27 @@ public class DefaultReactiveDataAccessStrategy implements ReactiveDataAccessStra /* * (non-Javadoc) - * @see org.springframework.data.r2dbc.core.ReactiveDataAccessStrategy#processNamedParameters(java.lang.String, java.util.Map) + * @see org.springframework.data.r2dbc.core.ReactiveDataAccessStrategy#processNamedParameters(java.lang.String, org.springframework.data.r2dbc.core.ReactiveDataAccessStrategy.NamedParameterProvider) */ @Override - public PreparedOperation processNamedParameters(String query, Map bindings) { - return this.expander.expand(query, this.dialect.getBindMarkersFactory(), new MapBindParameterSource(bindings)); + public PreparedOperation processNamedParameters(String query, NamedParameterProvider parameterProvider) { + + List parameterNames = this.expander.getParameterNames(query); + + Map namedBindings = new LinkedHashMap<>(parameterNames.size()); + for (String parameterName : parameterNames) { + + SettableValue value = parameterProvider.getParameter(parameterNames.indexOf(parameterName), parameterName); + + if (value == null) { + throw new InvalidDataAccessApiUsageException( + String.format("No parameter specified for [%s] in query [%s]", parameterName, query)); + } + + namedBindings.put(parameterName, value); + } + + return this.expander.expand(query, this.dialect.getBindMarkersFactory(), new MapBindParameterSource(namedBindings)); } /* diff --git a/src/main/java/org/springframework/data/r2dbc/core/NamedParameterExpander.java b/src/main/java/org/springframework/data/r2dbc/core/NamedParameterExpander.java index adcffec1..425babc9 100644 --- a/src/main/java/org/springframework/data/r2dbc/core/NamedParameterExpander.java +++ b/src/main/java/org/springframework/data/r2dbc/core/NamedParameterExpander.java @@ -16,6 +16,7 @@ package org.springframework.data.r2dbc.core; import java.util.LinkedHashMap; +import java.util.List; import java.util.Map; import org.apache.commons.logging.Log; @@ -137,4 +138,14 @@ public class NamedParameterExpander { return expanded; } + + /** + * Parse the SQL statement and locate any placeholders or named parameters. Named parameters are returned as result of + * this method invocation. + * + * @return the parameter names. + */ + public List getParameterNames(String sql) { + return getParsedSql(sql).getParameterNames(); + } } diff --git a/src/main/java/org/springframework/data/r2dbc/core/ReactiveDataAccessStrategy.java b/src/main/java/org/springframework/data/r2dbc/core/ReactiveDataAccessStrategy.java index 0d298f52..70deea30 100644 --- a/src/main/java/org/springframework/data/r2dbc/core/ReactiveDataAccessStrategy.java +++ b/src/main/java/org/springframework/data/r2dbc/core/ReactiveDataAccessStrategy.java @@ -19,12 +19,12 @@ import io.r2dbc.spi.Row; import io.r2dbc.spi.RowMetadata; import java.util.List; -import java.util.Map; import java.util.function.BiFunction; import org.springframework.data.r2dbc.convert.R2dbcConverter; import org.springframework.data.r2dbc.mapping.OutboundRow; import org.springframework.data.r2dbc.mapping.SettableValue; +import org.springframework.lang.Nullable; /** * Data access strategy that generalizes convenience operations using mapped entities. Typically used internally by @@ -66,13 +66,14 @@ public interface ReactiveDataAccessStrategy { String getTableName(Class type); /** - * Expand named parameters and return a {@link PreparedOperations} wrapping named bindings. - * + * Expand named parameters and return a {@link PreparedOperation} wrapping the given bindings. + * * @param query the query to expand. - * @param bindings named parameter bindings. - * @return the {@link PreparedOperation} encapsulating expanded SQL and bindings. + * @param parameterProvider indexed parameter bindings. + * @return the {@link PreparedOperation} encapsulating expanded SQL and namedBindings. + * @throws org.springframework.dao.InvalidDataAccessApiUsageException if a named parameter value cannot be resolved. */ - PreparedOperation processNamedParameters(String query, Map bindings); + PreparedOperation processNamedParameters(String query, NamedParameterProvider parameterProvider); /** * Returns the {@link org.springframework.data.r2dbc.dialect.R2dbcDialect}-specific {@link StatementMapper}. @@ -88,4 +89,22 @@ public interface ReactiveDataAccessStrategy { */ R2dbcConverter getConverter(); + /** + * Interface to retrieve parameters for named parameter processing. + */ + @FunctionalInterface + interface NamedParameterProvider { + + /** + * Returns the {@link SettableValue value} for a parameter identified either by name or by index. + * + * @param index parameter index according the parameter discovery order. + * @param name name of the parameter. + * @return the bindable value. Returning a {@literal null} value raises + * {@link org.springframework.dao.InvalidDataAccessApiUsageException} in named parameter processing. + */ + @Nullable + SettableValue getParameter(int index, String name); + } + } diff --git a/src/main/java/org/springframework/data/r2dbc/repository/query/StringBasedR2dbcQuery.java b/src/main/java/org/springframework/data/r2dbc/repository/query/StringBasedR2dbcQuery.java index 8840d840..a8799ed0 100644 --- a/src/main/java/org/springframework/data/r2dbc/repository/query/StringBasedR2dbcQuery.java +++ b/src/main/java/org/springframework/data/r2dbc/repository/query/StringBasedR2dbcQuery.java @@ -15,6 +15,11 @@ */ package org.springframework.data.r2dbc.repository.query; +import java.util.Map; +import java.util.Optional; +import java.util.concurrent.ConcurrentHashMap; +import java.util.regex.Pattern; + import org.springframework.data.r2dbc.convert.R2dbcConverter; import org.springframework.data.r2dbc.core.DatabaseClient; import org.springframework.data.r2dbc.core.DatabaseClient.BindSpec; @@ -35,6 +40,7 @@ import org.springframework.util.Assert; public class StringBasedR2dbcQuery extends AbstractR2dbcQuery { private final String sql; + private final Map namedParameters = new ConcurrentHashMap<>(); /** * Creates a new {@link StringBasedR2dbcQuery} for the given {@link StringBasedR2dbcQuery}, {@link DatabaseClient}, @@ -90,22 +96,48 @@ public class StringBasedR2dbcQuery extends AbstractR2dbcQuery { Parameters bindableParameters = accessor.getBindableParameters(); int index = 0; + int bindingIndex = 0; for (Object value : accessor.getValues()) { - Parameter bindableParameter = bindableParameters.getBindableParameter(index); + Parameter bindableParameter = bindableParameters.getBindableParameter(index++); - if (value == null) { - if (accessor.hasBindableNullValue()) { - bindSpecToUse = bindSpecToUse.bindNull(index++, bindableParameter.getType()); + Optional name = bindableParameter.getName(); + if (isNamedParameter(name)) { + if (value == null) { + if (accessor.hasBindableNullValue()) { + bindSpecToUse = bindSpecToUse.bindNull(name.get(), bindableParameter.getType()); + } + } else { + bindSpecToUse = bindSpecToUse.bind(name.get(), value); } } else { - bindSpecToUse = bindSpecToUse.bind(index++, value); + if (value == null) { + if (accessor.hasBindableNullValue()) { + bindSpecToUse = bindSpecToUse.bindNull(bindingIndex++, bindableParameter.getType()); + } + } else { + bindSpecToUse = bindSpecToUse.bind(bindingIndex++, value); + } } } return bindSpecToUse; } + private boolean isNamedParameter(Optional name) { + + if (!name.isPresent()) { + return false; + } + + return namedParameters.computeIfAbsent(name.get(), it -> { + + Pattern namedParameterPattern = Pattern.compile("(\\W)" + Pattern.quote(it) + "(\\W|$)"); + return namedParameterPattern.matcher(this.get()).find(); + }); + + } + @Override public String get() { return sql; diff --git a/src/test/java/org/springframework/data/r2dbc/core/AbstractDatabaseClientIntegrationTests.java b/src/test/java/org/springframework/data/r2dbc/core/AbstractDatabaseClientIntegrationTests.java index 16e3e692..f0f7d460 100644 --- a/src/test/java/org/springframework/data/r2dbc/core/AbstractDatabaseClientIntegrationTests.java +++ b/src/test/java/org/springframework/data/r2dbc/core/AbstractDatabaseClientIntegrationTests.java @@ -234,7 +234,7 @@ public abstract class AbstractDatabaseClientIntegrationTests extends R2dbcIntegr .using(legoSet) // .fetch() // .rowsUpdated() // - .then() + .then() // .as(StepVerifier::create) // .verifyComplete(); diff --git a/src/test/java/org/springframework/data/r2dbc/core/DefaultDatabaseClientUnitTests.java b/src/test/java/org/springframework/data/r2dbc/core/DefaultDatabaseClientUnitTests.java index f545b586..4f82b74e 100644 --- a/src/test/java/org/springframework/data/r2dbc/core/DefaultDatabaseClientUnitTests.java +++ b/src/test/java/org/springframework/data/r2dbc/core/DefaultDatabaseClientUnitTests.java @@ -26,6 +26,8 @@ import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; import reactor.test.StepVerifier; +import java.util.Arrays; + import org.junit.Before; import org.junit.Test; import org.junit.runner.RunWith; @@ -168,6 +170,31 @@ public class DefaultDatabaseClientUnitTests { verify(statement).bindNull(0, String.class); } + @Test // gh-178 + public void executeShouldBindNamedValuesFromIndexes() { + + Statement statement = mock(Statement.class); + when(connection.createStatement("SELECT id, name, manual FROM legoset WHERE name IN ($1, $2, $3)")) + .thenReturn(statement); + when(statement.execute()).thenReturn(Mono.empty()); + + DefaultDatabaseClient databaseClient = (DefaultDatabaseClient) DatabaseClient.builder() + .connectionFactory(connectionFactory) + .dataAccessStrategy(new DefaultReactiveDataAccessStrategy(PostgresDialect.INSTANCE)).build(); + + databaseClient.execute("SELECT id, name, manual FROM legoset WHERE name IN (:name)") // + .bind(0, Arrays.asList("unknown", "dunno", "other")) // + .then() // + .as(StepVerifier::create) // + .verifyComplete(); + + verify(statement).bind(0, "unknown"); + verify(statement).bind(1, "dunno"); + verify(statement).bind(2, "other"); + verify(statement).execute(); + verifyNoMoreInteractions(statement); + } + @Test // gh-128, gh-162 public void executeShouldBindValues() { 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 252f53ec..9e3aa62f 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 @@ -37,6 +37,7 @@ import org.springframework.data.repository.Repository; import org.springframework.data.repository.core.RepositoryMetadata; import org.springframework.data.repository.core.support.AbstractRepositoryMetadata; import org.springframework.data.repository.query.ExtensionAwareQueryMethodEvaluationContextProvider; +import org.springframework.data.repository.query.Param; import org.springframework.expression.spel.standard.SpelExpressionParser; import org.springframework.util.ReflectionUtils; @@ -67,6 +68,7 @@ public class StringBasedR2dbcQueryUnitTests { this.factory = new SpelAwareProxyProjectionFactory(); when(bindSpec.bind(anyInt(), any())).thenReturn(bindSpec); + when(bindSpec.bind(anyString(), any())).thenReturn(bindSpec); } @Test @@ -83,6 +85,48 @@ public class StringBasedR2dbcQueryUnitTests { verify(bindSpec).bind(0, "White"); } + @Test + public void bindsByNamedParameter() { + + StringBasedR2dbcQuery query = getQueryMethod("findByNamedParameter", String.class); + R2dbcParameterAccessor accessor = new R2dbcParameterAccessor(query.getQueryMethod(), "White"); + + BindableQuery stringQuery = query.createQuery(accessor); + + assertThat(stringQuery.get()).isEqualTo("SELECT * FROM person WHERE lastname = :lastname"); + assertThat(stringQuery.bind(bindSpec)).isNotNull(); + + verify(bindSpec).bind("lastname", "White"); + } + + @Test + public void bindsByBindmarker() { + + StringBasedR2dbcQuery query = getQueryMethod("findByNamedBindMarker", String.class); + R2dbcParameterAccessor accessor = new R2dbcParameterAccessor(query.getQueryMethod(), "White"); + + BindableQuery stringQuery = query.createQuery(accessor); + + assertThat(stringQuery.get()).isEqualTo("SELECT * FROM person WHERE lastname = @lastname"); + assertThat(stringQuery.bind(bindSpec)).isNotNull(); + + verify(bindSpec).bind("lastname", "White"); + } + + @Test + public void bindsByIndexWithNamedParameter() { + + StringBasedR2dbcQuery query = getQueryMethod("findNotByNamedBindMarker", String.class); + R2dbcParameterAccessor accessor = new R2dbcParameterAccessor(query.getQueryMethod(), "White"); + + BindableQuery stringQuery = query.createQuery(accessor); + + assertThat(stringQuery.get()).isEqualTo("SELECT * FROM person WHERE lastname = :unknown"); + assertThat(stringQuery.bind(bindSpec)).isNotNull(); + + verify(bindSpec).bind(0, "White"); + } + private StringBasedR2dbcQuery getQueryMethod(String name, Class... args) { Method method = ReflectionUtils.findMethod(SampleRepository.class, name, args); @@ -98,6 +142,15 @@ public class StringBasedR2dbcQueryUnitTests { @Query("SELECT * FROM person WHERE lastname = $1") Person findByLastname(String lastname); + + @Query("SELECT * FROM person WHERE lastname = :lastname") + Person findByNamedParameter(@Param("lastname") String lastname); + + @Query("SELECT * FROM person WHERE lastname = :unknown") + Person findNotByNamedBindMarker(String lastname); + + @Query("SELECT * FROM person WHERE lastname = @lastname") + Person findByNamedBindMarker(@Param("lastname") String lastname); } static class Person {