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 c90c8ec..0e87081 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 ef2e59c..5d18989 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 adcffec..425babc 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 0d298f5..70deea3 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 8840d84..a8799ed 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 16e3e69..f0f7d46 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 f545b58..4f82b74 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 252f53e..9e3aa62 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 {