diff --git a/src/main/java/org/springframework/data/r2dbc/core/NamedParameterUtils.java b/src/main/java/org/springframework/data/r2dbc/core/NamedParameterUtils.java index 8766230..8fdc79f 100644 --- a/src/main/java/org/springframework/data/r2dbc/core/NamedParameterUtils.java +++ b/src/main/java/org/springframework/data/r2dbc/core/NamedParameterUtils.java @@ -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 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 bindMarkers = getBindMarkers(identifier); + List> bindMarkers = getBindMarkers(identifier); if (bindMarkers == null) { @@ -506,28 +508,30 @@ abstract class NamedParameterUtils { return; } - if (value instanceof Collection) { - Collection collection = (Collection) value; + for (List outer : bindMarkers) { + if (value instanceof Collection) { + Collection collection = (Collection) value; - Iterator iterator = collection.iterator(); - Iterator markers = bindMarkers.iterator(); + Iterator iterator = collection.iterator(); + Iterator 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 bindMarkers = getBindMarkers(identifier); + List> bindMarkers = getBindMarkers(identifier); if (bindMarkers == null) { @@ -554,12 +558,15 @@ abstract class NamedParameterUtils { return; } - for (BindMarker bindMarker : bindMarkers) { - bindMarker.bindNull(target, valueType); + for (List outer : bindMarkers) { + for (BindMarker bindMarker : outer) { + bindMarker.bindNull(target, valueType); + } } } - List getBindMarkers(String identifier) { + @Nullable + List> getBindMarkers(String identifier) { List parameters = this.parameters.getMarker(identifier); @@ -567,10 +574,9 @@ abstract class NamedParameterUtils { return null; } - List markers = new ArrayList<>(); - + List> markers = new ArrayList<>(); for (NamedParameters.NamedParameter parameter : parameters) { - markers.addAll(parameter.placeholders); + markers.add(new ArrayList<>(parameter.placeholders)); } return markers; diff --git a/src/test/java/org/springframework/data/r2dbc/core/NamedParameterUtilsUnitTests.java b/src/test/java/org/springframework/data/r2dbc/core/NamedParameterUtilsUnitTests.java index 0168375..b10c320 100644 --- a/src/test/java/org/springframework/data/r2dbc/core/NamedParameterUtilsUnitTests.java +++ b/src/test/java/org/springframework/data/r2dbc/core/NamedParameterUtilsUnitTests.java @@ -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 operation = NamedParameterUtils + .substituteNamedParameters(parsedSql, BindMarkersFactory.anonymous("?"), new MapBindParameterSource( + Collections.singletonMap("names", SettableValue.from(Arrays.asList("1", "2", "3"))))); + + List 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 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 operation = NamedParameterUtils.substituteNamedParameters( + parsedSql, BindMarkersFactory.anonymous("?"), new MapBindParameterSource(parameterMap)); + + List 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 operation = NamedParameterUtils + .substituteNamedParameters(parsedSql, BindMarkersFactory.indexed("$", 1), new MapBindParameterSource( + Collections.singletonMap("names", SettableValue.from(Arrays.asList("1", "2", "3"))))); + + List 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 bindings; + + BindingCaptor(List 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(); }