From 22c236ca996c552509740f22e7bf45bc263d3fad Mon Sep 17 00:00:00 2001 From: Mark Paluch Date: Tue, 6 Feb 2024 14:53:34 +0100 Subject: [PATCH] Inspect collection elements during write value conversion. We now deeply introspect collection and map elements when obtaining write values to ensure proper UDT and tuple conversions. Previously, collections containing UDT values did a pass-thru of values instead of applying UDT mapping. Closes #1473 --- .../convert/MappingCassandraConverter.java | 21 +++++++++++++-- ...MappingCassandraConverterUDTUnitTests.java | 26 +++++++++++++++++++ .../StringBasedCassandraQueryUnitTests.java | 25 ++++++++++++++++++ 3 files changed, 70 insertions(+), 2 deletions(-) diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/convert/MappingCassandraConverter.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/convert/MappingCassandraConverter.java index bebabc131..5d124e7ff 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/convert/MappingCassandraConverter.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/convert/MappingCassandraConverter.java @@ -968,7 +968,13 @@ public class MappingCassandraConverter extends AbstractCassandraConverter ColumnType componentType = type.getRequiredComponentType(); for (Object element : source) { - converted.add(getWriteValue(element, componentType)); + + ColumnType elementType = componentType; + if (elementType.getType() == Object.class) { + elementType = getColumnTypeResolver().resolve(element); + } + + converted.add(getWriteValue(element, elementType)); } return converted; @@ -982,7 +988,18 @@ public class MappingCassandraConverter extends AbstractCassandraConverter ColumnType valueType = type.getRequiredMapValueType(); for (Entry entry : source.entrySet()) { - converted.put(getWriteValue(entry.getKey(), keyType), getWriteValue(entry.getValue(), valueType)); + + ColumnType elementKeyType = keyType; + if (elementKeyType.getType() == Object.class) { + elementKeyType = getColumnTypeResolver().resolve(entry.getKey()); + } + + ColumnType elementValueType = valueType; + if (elementValueType.getType() == Object.class) { + elementValueType = getColumnTypeResolver().resolve(entry.getValue()); + } + + converted.put(getWriteValue(entry.getKey(), elementKeyType), getWriteValue(entry.getValue(), elementValueType)); } return converted; diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/convert/MappingCassandraConverterUDTUnitTests.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/convert/MappingCassandraConverterUDTUnitTests.java index cde6297f5..b4afd1b2f 100755 --- a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/convert/MappingCassandraConverterUDTUnitTests.java +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/convert/MappingCassandraConverterUDTUnitTests.java @@ -25,6 +25,7 @@ import java.util.HashMap; import java.util.LinkedHashMap; import java.util.List; import java.util.Map; +import java.util.Set; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; @@ -199,6 +200,31 @@ class MappingCassandraConverterUDTUnitTests { + "VALUES ('foo',{zip:'69469',city:'Weinheim',streetlines:['Heckenpfad','14']})"); } + @Test // GH-1473 + void shouldWriteMapCorrectly() { + + Manufacturer manufacturer = new Manufacturer("foo", "bar"); + AddressUserType addressUserType = prepareAddressUserType(); + + Map value = Map.of(manufacturer, addressUserType); + Map writeValue = (Map) converter.convertToColumnType(value); + Map.Entry entry = writeValue.entrySet().iterator().next(); + + assertThat(entry.getKey()).isInstanceOf(UdtValue.class); + assertThat(entry.getValue()).isInstanceOf(UdtValue.class); + } + + @Test // GH-1473 + void shouldWriteSetCorrectly() { + + AddressUserType addressUserType = prepareAddressUserType(); + + Set value = Set.of(addressUserType); + Set writeValue = (Set) converter.convertToColumnType(value); + + assertThat(writeValue.iterator().next()).isInstanceOf(UdtValue.class); + } + private static AddressUserType prepareAddressUserType() { AddressUserType addressUserType = new AddressUserType(); diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/query/StringBasedCassandraQueryUnitTests.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/query/StringBasedCassandraQueryUnitTests.java index 8708634b1..71df5439a 100755 --- a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/query/StringBasedCassandraQueryUnitTests.java +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/query/StringBasedCassandraQueryUnitTests.java @@ -26,6 +26,7 @@ import java.time.LocalDate; import java.util.Arrays; import java.util.Collection; import java.util.HashSet; +import java.util.Set; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; @@ -352,6 +353,27 @@ class StringBasedCassandraQueryUnitTests { assertThat(actual.getPositionalValues().get(0)).isInstanceOf(UdtValue.class); } + @Test // GH-1473 + void bindsCollectionOfMappedUdtPropertyCorrectly() { + + UserDefinedType addressType = UserDefinedTypeBuilder.forName("address").withField("city", DataTypes.TEXT) + .withField("country", DataTypes.TEXT).build(); + + when(userTypeResolver.resolveType(CqlIdentifier.fromCql("address"))).thenReturn(addressType); + + StringBasedCassandraQuery cassandraQuery = getQueryMethod("findByMainAddress", Set.class); + CassandraParameterAccessor accessor = new ConvertingParameterAccessor(converter, + new CassandraParametersParameterAccessor(cassandraQuery.getQueryMethod(), Set.of(new AddressType()))); + + SimpleStatement actual = cassandraQuery.createQuery(accessor); + + assertThat(actual.getQuery()).isEqualTo("SELECT * FROM person WHERE address=?;"); + assertThat(actual.getPositionalValues().get(0)).isInstanceOf(Set.class); + + Set set = (Set) actual.getPositionalValues().get(0); + assertThat(set.iterator().next()).isInstanceOf(UdtValue.class); + } + @Test // DATACASS-172 void bindsUdtValuePropertyCorrectly() { @@ -467,6 +489,9 @@ class StringBasedCassandraQueryUnitTests { @Query("SELECT * FROM person WHERE address=?0;") Person findByMainAddress(AddressType address); + @Query("SELECT * FROM person WHERE address=?0;") + Person findByMainAddress(Set address); + @Query("SELECT * FROM person WHERE address=?0;") Person findByMainAddress(UdtValue UdtValue);