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