From 8893a4650b0fb25a2b6acb85fd611bfeebab7ce9 Mon Sep 17 00:00:00 2001 From: Mark Paluch Date: Thu, 20 Jan 2022 13:36:11 +0100 Subject: [PATCH] Guard access to identifier property. We now check in all places where we optionally use the Id property that an entity actually has an Id property and fall back leniently if the entity doesn't have an identifier property. Also, we use IdentifierAccessor consistently if the property is an identifier property. See #711 --- .../r2dbc/convert/MappingR2dbcConverter.java | 14 +++++++++- .../data/r2dbc/core/R2dbcEntityTemplate.java | 8 +++++- .../repository/query/R2dbcQueryCreator.java | 8 ++++-- .../MappingR2dbcConverterUnitTests.java | 28 +++++++++++++++++++ 4 files changed, 54 insertions(+), 4 deletions(-) diff --git a/src/main/java/org/springframework/data/r2dbc/convert/MappingR2dbcConverter.java b/src/main/java/org/springframework/data/r2dbc/convert/MappingR2dbcConverter.java index d8680ae..5e05c78 100644 --- a/src/main/java/org/springframework/data/r2dbc/convert/MappingR2dbcConverter.java +++ b/src/main/java/org/springframework/data/r2dbc/convert/MappingR2dbcConverter.java @@ -31,6 +31,7 @@ import org.springframework.core.CollectionFactory; import org.springframework.core.convert.ConversionService; import org.springframework.dao.InvalidDataAccessApiUsageException; import org.springframework.data.convert.CustomConversions; +import org.springframework.data.mapping.IdentifierAccessor; import org.springframework.data.mapping.MappingException; import org.springframework.data.mapping.PersistentProperty; import org.springframework.data.mapping.PersistentPropertyAccessor; @@ -370,7 +371,14 @@ public class MappingR2dbcConverter extends BasicRelationalConverter implements R continue; } - Object value = accessor.getProperty(property); + Object value; + + if (property.isIdProperty()) { + IdentifierAccessor identifierAccessor = entity.getIdentifierAccessor(accessor.getBean()); + value = identifierAccessor.getIdentifier(); + } else { + value = accessor.getProperty(property); + } if (value == null) { writeNullInternal(sink, property); @@ -600,6 +608,10 @@ public class MappingR2dbcConverter extends BasicRelationalConverter implements R Class userClass = ClassUtils.getUserClass(object); RelationalPersistentEntity entity = getMappingContext().getRequiredPersistentEntity(userClass); + if (!entity.hasIdProperty()) { + return (row, rowMetadata) -> object; + } + return (row, metadata) -> { PersistentPropertyAccessor propertyAccessor = entity.getPropertyAccessor(object); diff --git a/src/main/java/org/springframework/data/r2dbc/core/R2dbcEntityTemplate.java b/src/main/java/org/springframework/data/r2dbc/core/R2dbcEntityTemplate.java index f23fd9f..0f2c515 100644 --- a/src/main/java/org/springframework/data/r2dbc/core/R2dbcEntityTemplate.java +++ b/src/main/java/org/springframework/data/r2dbc/core/R2dbcEntityTemplate.java @@ -65,6 +65,7 @@ import org.springframework.data.relational.core.query.CriteriaDefinition; import org.springframework.data.relational.core.query.Query; import org.springframework.data.relational.core.query.Update; import org.springframework.data.relational.core.sql.Expression; +import org.springframework.data.relational.core.sql.Expressions; import org.springframework.data.relational.core.sql.Functions; import org.springframework.data.relational.core.sql.SqlIdentifier; import org.springframework.data.relational.core.sql.Table; @@ -322,7 +323,11 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw StatementMapper.SelectSpec selectSpec = statementMapper // .createSelect(tableName) // .doWithTable((table, spec) -> { - return spec.withProjection(Functions.count(table.column(entity.getRequiredIdProperty().getColumnName()))); + + Expression countExpression = entity.hasIdProperty() + ? table.column(entity.getRequiredIdProperty().getColumnName()) + : Expressions.asterisk(); + return spec.withProjection(Functions.count(countExpression)); }); Optional criteria = query.getCriteria(); @@ -834,6 +839,7 @@ public class R2dbcEntityTemplate implements R2dbcEntityOperations, BeanFactoryAw } private Query getByIdQuery(T entity, RelationalPersistentEntity persistentEntity) { + if (!persistentEntity.hasIdProperty()) { throw new MappingException("No id property found for object of type " + persistentEntity.getType() + "!"); } 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 6164f09..3abc8c2 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 @@ -29,6 +29,7 @@ import org.springframework.data.relational.core.mapping.RelationalPersistentProp import org.springframework.data.relational.core.query.Criteria; import org.springframework.data.relational.core.sql.Column; import org.springframework.data.relational.core.sql.Expression; +import org.springframework.data.relational.core.sql.Expressions; import org.springframework.data.relational.core.sql.Functions; import org.springframework.data.relational.core.sql.SqlIdentifier; import org.springframework.data.relational.core.sql.Table; @@ -164,8 +165,11 @@ class R2dbcQueryCreator extends RelationalQueryCreator> { .collect(Collectors.toList()); } else if (tree.isCountProjection()) { - SqlIdentifier idColumn = entityMetadata.getTableEntity().getRequiredIdProperty().getColumnName(); - expressions = Collections.singletonList(Functions.count(table.column(idColumn))); + Expression countExpression = entityMetadata.getTableEntity().hasIdProperty() + ? table.column(entityMetadata.getTableEntity().getRequiredIdProperty().getColumnName()) + : Expressions.asterisk(); + + expressions = Collections.singletonList(Functions.count(countExpression)); } else { expressions = dataAccessStrategy.getAllColumns(entityToRead).stream() .map(table::column) diff --git a/src/test/java/org/springframework/data/r2dbc/convert/MappingR2dbcConverterUnitTests.java b/src/test/java/org/springframework/data/r2dbc/convert/MappingR2dbcConverterUnitTests.java index ef253b2..05d7f59 100644 --- a/src/test/java/org/springframework/data/r2dbc/convert/MappingR2dbcConverterUnitTests.java +++ b/src/test/java/org/springframework/data/r2dbc/convert/MappingR2dbcConverterUnitTests.java @@ -43,6 +43,7 @@ import org.springframework.data.annotation.Id; import org.springframework.data.annotation.Transient; import org.springframework.data.convert.ReadingConverter; import org.springframework.data.convert.WritingConverter; +import org.springframework.data.domain.Persistable; import org.springframework.data.r2dbc.dialect.PostgresDialect; import org.springframework.data.r2dbc.mapping.OutboundRow; import org.springframework.data.r2dbc.mapping.R2dbcMappingContext; @@ -250,6 +251,18 @@ public class MappingR2dbcConverterUnitTests { assertThat(result.person).isNull(); } + @Test // GH-711 + void writeShouldObtainIdFromIdentifierAccessor() { + + PersistableEntity entity = new PersistableEntity(); + entity.id = null; + + OutboundRow row = new OutboundRow(); + converter.write(entity, row); + + assertThat(row).containsEntry(SqlIdentifier.unquoted("id"), Parameter.from(42L)); + } + @AllArgsConstructor static class Person { @Id String id; @@ -406,4 +419,19 @@ public class MappingR2dbcConverterUnitTests { this.name = name; } } + + static class PersistableEntity implements Persistable { + + @Id String id; + + @Override + public Long getId() { + return 42L; + } + + @Override + public boolean isNew() { + return false; + } + } }