diff --git a/src/main/java/org/springframework/data/r2dbc/repository/query/R2dbcQueryCreator.java b/src/main/java/org/springframework/data/r2dbc/repository/query/R2dbcQueryCreator.java index 45ec1ac1..397fc9c6 100644 --- a/src/main/java/org/springframework/data/r2dbc/repository/query/R2dbcQueryCreator.java +++ b/src/main/java/org/springframework/data/r2dbc/repository/query/R2dbcQueryCreator.java @@ -16,6 +16,7 @@ package org.springframework.data.r2dbc.repository.query; import java.util.ArrayList; +import java.util.Collections; import java.util.List; import java.util.stream.Collectors; @@ -27,7 +28,7 @@ import org.springframework.data.r2dbc.core.StatementMapper; import org.springframework.data.relational.core.mapping.RelationalPersistentEntity; import org.springframework.data.relational.core.mapping.RelationalPersistentProperty; import org.springframework.data.relational.core.query.Criteria; -import org.springframework.data.relational.core.sql.SqlIdentifier; +import org.springframework.data.relational.core.sql.*; import org.springframework.data.relational.repository.query.RelationalEntityMetadata; import org.springframework.data.relational.repository.query.RelationalParameterAccessor; import org.springframework.data.relational.repository.query.RelationalQueryCreator; @@ -41,6 +42,7 @@ import org.springframework.lang.Nullable; * @author Roman Chigvintsev * @author Mark Paluch * @author Mingyuan Wu + * @author Myeonghyeon Lee * @since 1.1 */ class R2dbcQueryCreator extends RelationalQueryCreator> { @@ -132,28 +134,39 @@ class R2dbcQueryCreator extends RelationalQueryCreator> { return statementMapper.getMappedObject(selectSpec); } - private SqlIdentifier[] getSelectProjection() { + private Expression[] getSelectProjection() { - List columnNames; + List expressions; + Table table = Table.create(entityMetadata.getTableName()); if (!projectedProperties.isEmpty()) { RelationalPersistentEntity entity = entityMetadata.getTableEntity(); - columnNames = new ArrayList<>(projectedProperties.size()); + expressions = new ArrayList<>(projectedProperties.size()); for (String projectedProperty : projectedProperties) { RelationalPersistentProperty property = entity.getPersistentProperty(projectedProperty); - columnNames.add(property != null ? property.getColumnName() : SqlIdentifier.unquoted(projectedProperty)); + Column column = table.column(property != null ? property.getColumnName() : SqlIdentifier.unquoted(projectedProperty)); + expressions.add(column); } } else if (tree.isExistsProjection()) { - columnNames = dataAccessStrategy.getIdentifierColumns(entityMetadata.getJavaType()); + + expressions = dataAccessStrategy.getIdentifierColumns(entityMetadata.getJavaType()).stream() + .map(table::column) + .collect(Collectors.toList()); + } else if (tree.isCountProjection()) { + + SqlIdentifier idColumn = entityMetadata.getTableEntity().getRequiredIdProperty().getColumnName(); + expressions = Collections.singletonList(Functions.count(table.column(idColumn))); } else { - columnNames = dataAccessStrategy.getAllColumns(entityMetadata.getJavaType()); + expressions = dataAccessStrategy.getAllColumns(entityMetadata.getJavaType()).stream() + .map(table::column) + .collect(Collectors.toList()); } - return columnNames.toArray(new SqlIdentifier[0]); + return expressions.toArray(new Expression[0]); } private Sort getSort(Sort sort) { diff --git a/src/test/java/org/springframework/data/r2dbc/repository/query/PartTreeR2dbcQueryUnitTests.java b/src/test/java/org/springframework/data/r2dbc/repository/query/PartTreeR2dbcQueryUnitTests.java index 927f0b97..f6ab7774 100644 --- a/src/test/java/org/springframework/data/r2dbc/repository/query/PartTreeR2dbcQueryUnitTests.java +++ b/src/test/java/org/springframework/data/r2dbc/repository/query/PartTreeR2dbcQueryUnitTests.java @@ -56,6 +56,7 @@ import org.springframework.data.repository.core.support.DefaultRepositoryMetadat * @author Roman Chigvintsev * @author Mark Paluch * @author Mingyuan Wu + * @author Myeonghyeon Lee */ @RunWith(MockitoJUnitRunner.class) public class PartTreeR2dbcQueryUnitTests { @@ -618,6 +619,18 @@ public class PartTreeR2dbcQueryUnitTests { + ".foo FROM " + TABLE + " WHERE " + TABLE + ".first_name = $1"); } + @Test // DATAJDBC-534 + public void createsQueryForCountProjection() throws Exception { + + R2dbcQueryMethod queryMethod = getQueryMethod("countByFirstName", String.class); + PartTreeR2dbcQuery r2dbcQuery = new PartTreeR2dbcQuery(queryMethod, databaseClient, r2dbcConverter, + dataAccessStrategy); + BindableQuery query = r2dbcQuery.createQuery((getAccessor(queryMethod, new Object[] { "John" }))); + + assertThat(query.get()) + .isEqualTo("SELECT COUNT(users.id) FROM " + TABLE + " WHERE " + TABLE + ".first_name = $1"); + } + private R2dbcQueryMethod getQueryMethod(String methodName, Class... parameterTypes) throws Exception { Method method = UserRepository.class.getMethod(methodName, parameterTypes); return new R2dbcQueryMethod(method, new DefaultRepositoryMetadata(UserRepository.class), @@ -699,6 +712,8 @@ public class PartTreeR2dbcQueryUnitTests { Mono findDistinctByFirstName(String firstName); Mono deleteByFirstName(String firstName); + + Mono countByFirstName(String firstName); } @Table("users")