From eb6303605c86966cd442b31e17ba0ea1922800ec Mon Sep 17 00:00:00 2001 From: Mark Paluch Date: Wed, 6 Dec 2023 15:29:12 +0100 Subject: [PATCH] Fix interface projection for entities that implement the interface. Using as with an interface that is implemented by the entity, we no longer attempt to instantiate the interface bur use the entity type instead. Closes #1690 --- .../data/r2dbc/core/R2dbcEntityTemplate.java | 62 ++++++++++++++----- .../core/R2dbcEntityTemplateUnitTests.java | 40 +++++++++++- .../ReactiveSelectOperationUnitTests.java | 4 +- 3 files changed, 88 insertions(+), 18 deletions(-) diff --git a/spring-data-r2dbc/src/main/java/org/springframework/data/r2dbc/core/R2dbcEntityTemplate.java b/spring-data-r2dbc/src/main/java/org/springframework/data/r2dbc/core/R2dbcEntityTemplate.java index aebcf866..431a9293 100644 --- a/spring-data-r2dbc/src/main/java/org/springframework/data/r2dbc/core/R2dbcEntityTemplate.java +++ b/spring-data-r2dbc/src/main/java/org/springframework/data/r2dbc/core/R2dbcEntityTemplate.java @@ -21,16 +21,18 @@ import io.r2dbc.spi.RowMetadata; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; -import java.beans.FeatureDescriptor; import java.util.Collections; +import java.util.LinkedHashSet; import java.util.List; import java.util.Map; import java.util.Optional; +import java.util.Set; import java.util.function.BiFunction; import java.util.function.Function; import java.util.stream.Collectors; import org.reactivestreams.Publisher; + import org.springframework.beans.BeansException; import org.springframework.beans.factory.BeanFactory; import org.springframework.beans.factory.BeanFactoryAware; @@ -46,7 +48,6 @@ import org.springframework.data.mapping.PersistentPropertyAccessor; import org.springframework.data.mapping.callback.ReactiveEntityCallbacks; import org.springframework.data.mapping.context.MappingContext; import org.springframework.data.projection.EntityProjection; -import org.springframework.data.projection.ProjectionInformation; import org.springframework.data.projection.SpelAwareProxyProjectionFactory; import org.springframework.data.r2dbc.convert.R2dbcConverter; import org.springframework.data.r2dbc.dialect.DialectResolver; @@ -56,6 +57,7 @@ import org.springframework.data.r2dbc.mapping.event.AfterConvertCallback; import org.springframework.data.r2dbc.mapping.event.AfterSaveCallback; import org.springframework.data.r2dbc.mapping.event.BeforeConvertCallback; import org.springframework.data.r2dbc.mapping.event.BeforeSaveCallback; +import org.springframework.data.relational.core.mapping.PersistentPropertyTranslator; import org.springframework.data.relational.core.mapping.RelationalPersistentEntity; import org.springframework.data.relational.core.mapping.RelationalPersistentProperty; import org.springframework.data.relational.core.query.Criteria; @@ -68,6 +70,7 @@ import org.springframework.data.relational.core.sql.Functions; import org.springframework.data.relational.core.sql.SqlIdentifier; import org.springframework.data.relational.core.sql.Table; import org.springframework.data.relational.domain.RowDocument; +import org.springframework.data.util.Predicates; import org.springframework.data.util.ProxyUtils; import org.springframework.lang.Nullable; import org.springframework.r2dbc.core.DatabaseClient; @@ -332,7 +335,7 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw StatementMapper.SelectSpec selectSpec = statementMapper // .createSelect(tableName) // - .doWithTable((table, spec) -> spec.withProjection(getSelectProjection(table, query, returnType))); + .doWithTable((table, spec) -> spec.withProjection(getSelectProjection(table, query, entityType, returnType))); if (query.getLimit() > 0) { selectSpec = selectSpec.limit(query.getLimit()); @@ -423,7 +426,8 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw } @Override - public RowsFetchSpec query(PreparedOperation operation, Class entityClass, Class resultType) throws DataAccessException { + public RowsFetchSpec query(PreparedOperation operation, Class entityClass, Class resultType) + throws DataAccessException { Assert.notNull(operation, "PreparedOperation must not be null"); Assert.notNull(entityClass, "Entity class must not be null"); @@ -759,18 +763,16 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw return (RelationalPersistentEntity) getRequiredEntity(entityType); } - private List getSelectProjection(Table table, Query query, Class returnType) { + private List getSelectProjection(Table table, Query query, Class entityType, Class returnType) { if (query.getColumns().isEmpty()) { - if (returnType.isInterface()) { + EntityProjection projection = converter.introspectProjection(returnType, entityType); - ProjectionInformation projectionInformation = projectionFactory.getProjectionInformation(returnType); + if (projection.isProjection() && projection.isClosedProjection()) { + + return computeProjectedFields(table, returnType, projection); - if (projectionInformation.isClosed()) { - return projectionInformation.getInputProperties().stream().map(FeatureDescriptor::getName).map(table::column) - .collect(Collectors.toList()); - } } return Collections.singletonList(table.asterisk()); @@ -779,6 +781,36 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw return query.getColumns().stream().map(table::column).collect(Collectors.toList()); } + @SuppressWarnings("unchecked") + private List computeProjectedFields(Table table, Class returnType, + EntityProjection projection) { + + if (returnType.isInterface()) { + + Set properties = new LinkedHashSet<>(); + projection.forEach(it -> { + properties.add(it.getPropertyPath().getSegment()); + }); + + return properties.stream().map(table::column).collect(Collectors.toList()); + } + + Set properties = new LinkedHashSet<>(); + // DTO projections use merged metadata between domain type and result type + PersistentPropertyTranslator translator = PersistentPropertyTranslator.create( + mappingContext.getRequiredPersistentEntity(projection.getDomainType()), + Predicates.negate(RelationalPersistentProperty::hasExplicitColumnName)); + + RelationalPersistentEntity persistentEntity = mappingContext + .getRequiredPersistentEntity(projection.getMappedType()); + for (RelationalPersistentProperty property : persistentEntity) { + properties.add(translator.translate(property).getColumnName()); + } + + return properties.stream().map(table::column).collect(Collectors.toList()); + } + + @SuppressWarnings("unchecked") public RowsFetchSpec getRowsFetchSpec(DatabaseClient.GenericExecuteSpec executeSpec, Class entityType, Class resultType) { @@ -791,13 +823,13 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw } else { EntityProjection projection = converter.introspectProjection(resultType, entityType); + Class typeToRead = projection.isProjection() ? resultType + : resultType.isInterface() ? (Class) entityType : resultType; rowMapper = (row, rowMetadata) -> { - RowDocument document = dataAccessStrategy.toRowDocument(resultType, row, rowMetadata.getColumnMetadatas()); - - return projection.isProjection() ? converter.project(projection, document) - : converter.read(resultType, document); + RowDocument document = dataAccessStrategy.toRowDocument(typeToRead, row, rowMetadata.getColumnMetadatas()); + return converter.project(projection, document); }; } diff --git a/spring-data-r2dbc/src/test/java/org/springframework/data/r2dbc/core/R2dbcEntityTemplateUnitTests.java b/spring-data-r2dbc/src/test/java/org/springframework/data/r2dbc/core/R2dbcEntityTemplateUnitTests.java index 1bd4a9e8..3097b840 100644 --- a/spring-data-r2dbc/src/test/java/org/springframework/data/r2dbc/core/R2dbcEntityTemplateUnitTests.java +++ b/spring-data-r2dbc/src/test/java/org/springframework/data/r2dbc/core/R2dbcEntityTemplateUnitTests.java @@ -121,6 +121,35 @@ public class R2dbcEntityTemplateUnitTests { .all() // .as(StepVerifier::create) // .assertNext(actual -> assertThat(actual.getName()).isEqualTo("Walter")).verifyComplete(); + + StatementRecorder.RecordedStatement statement = recorder.getCreatedStatement(s -> s.startsWith("SELECT")); + assertThat(statement.getSql()).isEqualTo("SELECT foo.THE_NAME FROM foo WHERE foo.THE_NAME = $1"); + } + + @Test // GH-1690 + void shouldProjectEntityUsingInheritedInterface() { + + MockRowMetadata metadata = MockRowMetadata.builder() + .columnMetadata(MockColumnMetadata.builder().name("THE_NAME").type(R2dbcType.VARCHAR).build()).build(); + MockResult result = MockResult.builder() + .row(MockRow.builder().identified("THE_NAME", Object.class, "Walter").metadata(metadata).build()).build(); + + recorder.addStubbing(s -> s.startsWith("SELECT"), result); + + entityTemplate.select(Person.class) // + .from("foo") // + .as(Named.class) // + .matching(Query.query(Criteria.where("name").is("Walter"))) // + .all() // + .as(StepVerifier::create) // + .assertNext(actual -> { + assertThat(actual.getName()).isEqualTo("Walter"); + assertThat(actual).isInstanceOf(Person.class); + }).verifyComplete(); + + + StatementRecorder.RecordedStatement statement = recorder.getCreatedStatement(s -> s.startsWith("SELECT")); + assertThat(statement.getSql()).isEqualTo("SELECT foo.* FROM foo WHERE foo.THE_NAME = $1"); } @Test // gh-469 @@ -558,11 +587,15 @@ public class R2dbcEntityTemplateUnitTests { record WithoutId(String name) { } + interface Named{ + String getName(); + } + record Person(@Id String id, @Column("THE_NAME") String name, - String description) { + String description) implements Named { public static Person empty() { return new Person(null, null, null); @@ -579,6 +612,11 @@ public class R2dbcEntityTemplateUnitTests { public Person withDescription(String description) { return this.description == description ? this : new Person(this.id, this.name, description); } + + @Override + public String getName() { + return name(); + } } interface PersonProjection { diff --git a/spring-data-r2dbc/src/test/java/org/springframework/data/r2dbc/core/ReactiveSelectOperationUnitTests.java b/spring-data-r2dbc/src/test/java/org/springframework/data/r2dbc/core/ReactiveSelectOperationUnitTests.java index 3f35c8b6..2a2bfa55 100644 --- a/spring-data-r2dbc/src/test/java/org/springframework/data/r2dbc/core/ReactiveSelectOperationUnitTests.java +++ b/spring-data-r2dbc/src/test/java/org/springframework/data/r2dbc/core/ReactiveSelectOperationUnitTests.java @@ -102,7 +102,7 @@ public class ReactiveSelectOperationUnitTests { assertThat(statement.getSql()).isEqualTo("SELECT person.THE_NAME FROM person WHERE person.THE_NAME = $1"); } - @Test // gh-220 + @Test // GH-220, GH-1690 void shouldSelectAsWithColumnName() { MockRowMetadata metadata = MockRowMetadata.builder() @@ -123,7 +123,7 @@ public class ReactiveSelectOperationUnitTests { StatementRecorder.RecordedStatement statement = recorder.getCreatedStatement(s -> s.startsWith("SELECT")); - assertThat(statement.getSql()).isEqualTo("SELECT person.* FROM person WHERE person.THE_NAME = $1"); + assertThat(statement.getSql()).isEqualTo("SELECT person.id, person.a_different_name FROM person WHERE person.THE_NAME = $1"); } @Test // gh-220