Fix parameter binding of reused named parameters using anonymous bind markers.

We now properly bind values for reused named parameters correctly when using anonymous bind markers (e.g. ? for MySQL). Previously, the subsequent usages of named parameters especially with IN parameters were left not bound.

Closes #778
This commit is contained in:
Mark Paluch
2022-08-16 11:11:33 +02:00
parent b30a9d4b5e
commit 39216db553
2 changed files with 113 additions and 24 deletions

View File

@@ -27,6 +27,7 @@ import java.util.TreeMap;
import org.springframework.dao.InvalidDataAccessApiUsageException;
import org.springframework.data.r2dbc.dialect.BindTarget;
import org.springframework.lang.Nullable;
import org.springframework.r2dbc.core.binding.BindMarker;
import org.springframework.r2dbc.core.binding.BindMarkers;
import org.springframework.r2dbc.core.binding.BindMarkersFactory;
@@ -435,6 +436,7 @@ abstract class NamedParameterUtils {
return param;
}
@Nullable
List<NamedParameter> getMarker(String name) {
return this.references.get(name);
}
@@ -498,7 +500,7 @@ abstract class NamedParameterUtils {
@SuppressWarnings("unchecked")
public void bind(org.springframework.r2dbc.core.binding.BindTarget target, String identifier, Object value) {
List<BindMarker> bindMarkers = getBindMarkers(identifier);
List<List<BindMarker>> bindMarkers = getBindMarkers(identifier);
if (bindMarkers == null) {
@@ -506,28 +508,30 @@ abstract class NamedParameterUtils {
return;
}
if (value instanceof Collection) {
Collection<Object> collection = (Collection<Object>) value;
for (List<BindMarker> outer : bindMarkers) {
if (value instanceof Collection) {
Collection<Object> collection = (Collection<Object>) value;
Iterator<Object> iterator = collection.iterator();
Iterator<BindMarker> markers = bindMarkers.iterator();
Iterator<Object> iterator = collection.iterator();
Iterator<BindMarker> markers = outer.iterator();
while (iterator.hasNext()) {
while (iterator.hasNext()) {
Object valueToBind = iterator.next();
Object valueToBind = iterator.next();
if (valueToBind instanceof Object[]) {
Object[] objects = (Object[]) valueToBind;
for (Object object : objects) {
bind(target, markers, object);
if (valueToBind instanceof Object[]) {
Object[] objects = (Object[]) valueToBind;
for (Object object : objects) {
bind(target, markers, object);
}
} else {
bind(target, markers, valueToBind);
}
} else {
bind(target, markers, valueToBind);
}
}
} else {
for (BindMarker bindMarker : bindMarkers) {
bindMarker.bind(target, value);
} else {
for (BindMarker bindMarker : outer) {
bindMarker.bind(target, value);
}
}
}
}
@@ -546,7 +550,7 @@ abstract class NamedParameterUtils {
public void bindNull(org.springframework.r2dbc.core.binding.BindTarget target, String identifier,
Class<?> valueType) {
List<BindMarker> bindMarkers = getBindMarkers(identifier);
List<List<BindMarker>> bindMarkers = getBindMarkers(identifier);
if (bindMarkers == null) {
@@ -554,12 +558,15 @@ abstract class NamedParameterUtils {
return;
}
for (BindMarker bindMarker : bindMarkers) {
bindMarker.bindNull(target, valueType);
for (List<BindMarker> outer : bindMarkers) {
for (BindMarker bindMarker : outer) {
bindMarker.bindNull(target, valueType);
}
}
}
List<BindMarker> getBindMarkers(String identifier) {
@Nullable
List<List<BindMarker>> getBindMarkers(String identifier) {
List<NamedParameters.NamedParameter> parameters = this.parameters.getMarker(identifier);
@@ -567,10 +574,9 @@ abstract class NamedParameterUtils {
return null;
}
List<BindMarker> markers = new ArrayList<>();
List<List<BindMarker>> markers = new ArrayList<>();
for (NamedParameters.NamedParameter parameter : parameters) {
markers.addAll(parameter.placeholders);
markers.add(new ArrayList<>(parameter.placeholders));
}
return markers;

View File

@@ -18,10 +18,12 @@ package org.springframework.data.r2dbc.core;
import static org.assertj.core.api.Assertions.*;
import static org.mockito.Mockito.*;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.Collections;
import java.util.HashMap;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
@@ -458,6 +460,87 @@ class NamedParameterUtilsUnitTests {
});
}
@Test // GH-1306
void inCollectionSameParameterNameShouldBindAllAnonymousParameters() {
ParsedSql parsedSql = NamedParameterUtils.parseSqlStatement("select :names AND :names");
org.springframework.r2dbc.core.PreparedOperation<String> operation = NamedParameterUtils
.substituteNamedParameters(parsedSql, BindMarkersFactory.anonymous("?"), new MapBindParameterSource(
Collections.singletonMap("names", SettableValue.from(Arrays.asList("1", "2", "3")))));
List<String> bindings = new ArrayList<>();
operation.bindTo(new BindingCaptor(bindings));
assertThat(operation.get()).isEqualTo("select ?, ?, ? AND ?, ?, ?");
assertThat(bindings).contains("0: 1", "1: 2", "2: 3", "3: 1", "4: 2", "5: 3");
}
@Test // GH-1306
void complexInCollectionSameParameterNameShouldBindAllAnonymousParameters() {
Map<String, SettableValue> parameterMap = new HashMap<>();
parameterMap.put("names", SettableValue.from(Arrays.asList("1", "2", "3")));
parameterMap.put("hello", SettableValue.from("world"));
ParsedSql parsedSql = NamedParameterUtils.parseSqlStatement("select :names AND :hello OR :names");
org.springframework.r2dbc.core.PreparedOperation<String> operation = NamedParameterUtils.substituteNamedParameters(
parsedSql, BindMarkersFactory.anonymous("?"), new MapBindParameterSource(parameterMap));
List<String> bindings = new ArrayList<>();
operation.bindTo(new BindingCaptor(bindings));
assertThat(operation.get()).isEqualTo("select ?, ?, ? AND ? OR ?, ?, ?");
assertThat(bindings).contains("0: 1", "1: 2", "2: 3", "3: world", "4: 1", "5: 2", "6: 3");
}
@Test // GH-1306
void inCollectionSameParameterNameShouldBindAllNamedParameters() {
ParsedSql parsedSql = NamedParameterUtils.parseSqlStatement("select :names AND :names");
org.springframework.r2dbc.core.PreparedOperation<String> operation = NamedParameterUtils
.substituteNamedParameters(parsedSql, BindMarkersFactory.indexed("$", 1), new MapBindParameterSource(
Collections.singletonMap("names", SettableValue.from(Arrays.asList("1", "2", "3")))));
List<String> bindings = new ArrayList<>();
operation.bindTo(new BindingCaptor(bindings));
assertThat(operation.get()).isEqualTo("select $1, $2, $3 AND $1, $2, $3");
assertThat(bindings).containsOnly("0: 1", "1: 2", "2: 3");
}
static class BindingCaptor implements org.springframework.r2dbc.core.binding.BindTarget {
private final List<String> bindings;
BindingCaptor(List<String> bindings) {
this.bindings = bindings;
}
@Override
public void bind(String identifier, Object value) {
bindings.add(identifier + ": " + value);
}
@Override
public void bind(int index, Object value) {
bindings.add(index + ": " + value);
}
@Override
public void bindNull(String identifier, Class<?> type) {
}
@Override
public void bindNull(int index, Class<?> type) {
}
}
private String expand(ParsedSql sql) {
return NamedParameterUtils.substituteNamedParameters(sql, BIND_MARKERS, new MapBindParameterSource()).toQuery();
}