#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.
This commit is contained in:
@@ -345,7 +345,22 @@ class DefaultDatabaseClient implements DatabaseClient, ConnectionAccessor {
|
||||
|
||||
if (namedParameters) {
|
||||
|
||||
PreparedOperation<?> operation = dataAccessStrategy.processNamedParameters(sql, this.byName);
|
||||
Map<String, SettableValue> remainderByName = new LinkedHashMap<>(this.byName);
|
||||
Map<Integer, SettableValue> 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;
|
||||
}
|
||||
|
||||
@@ -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<String, SettableValue> bindings) {
|
||||
return this.expander.expand(query, this.dialect.getBindMarkersFactory(), new MapBindParameterSource(bindings));
|
||||
public PreparedOperation<?> processNamedParameters(String query, NamedParameterProvider parameterProvider) {
|
||||
|
||||
List<String> parameterNames = this.expander.getParameterNames(query);
|
||||
|
||||
Map<String, SettableValue> 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));
|
||||
}
|
||||
|
||||
/*
|
||||
|
||||
@@ -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<String> getParameterNames(String sql) {
|
||||
return getParsedSql(sql).getParameterNames();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<String, SettableValue> 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);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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<String, Boolean> 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<String> 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<String> 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;
|
||||
|
||||
@@ -234,7 +234,7 @@ public abstract class AbstractDatabaseClientIntegrationTests extends R2dbcIntegr
|
||||
.using(legoSet) //
|
||||
.fetch() //
|
||||
.rowsUpdated() //
|
||||
.then()
|
||||
.then() //
|
||||
.as(StepVerifier::create) //
|
||||
.verifyComplete();
|
||||
|
||||
|
||||
@@ -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() {
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user