diff --git a/pom.xml b/pom.xml index 8f9de8fe..4ada1017 100644 --- a/pom.xml +++ b/pom.xml @@ -67,6 +67,11 @@ spring-beans + + org.springframework + spring-jdbc + + org.springframework spring-core @@ -84,9 +89,12 @@ 2.2.8 test + - org.springframework - spring-jdbc + org.assertj + assertj-core + 3.6.2 + test diff --git a/src/main/java/org/springframework/data/jdbc/mapping/model/JdbcPersistentEntity.java b/src/main/java/org/springframework/data/jdbc/mapping/model/JdbcPersistentEntity.java index 7e91a166..217d8d70 100644 --- a/src/main/java/org/springframework/data/jdbc/mapping/model/JdbcPersistentEntity.java +++ b/src/main/java/org/springframework/data/jdbc/mapping/model/JdbcPersistentEntity.java @@ -19,11 +19,35 @@ import org.springframework.data.mapping.model.BasicPersistentEntity; import org.springframework.data.util.TypeInformation; /** + * meta data a repository might need for implementing persistence operations for instances of type {@code T} * @author Jens Schauder */ public class JdbcPersistentEntity extends BasicPersistentEntity { + private String tableName; + private String idColumn; + public JdbcPersistentEntity(TypeInformation information) { super(information); } + + public String getTableName() { + + if (tableName == null) + tableName = getType().getSimpleName(); + + return tableName; + } + + public String getIdColumn() { + + if (idColumn == null) + idColumn = getIdProperty().getName(); + + return idColumn; + } + + public Object getIdValue(T instance) { + return getPropertyAccessor(instance).getProperty(getIdProperty()); + } } diff --git a/src/main/java/org/springframework/data/jdbc/mapping/model/JdbcPersistentProperty.java b/src/main/java/org/springframework/data/jdbc/mapping/model/JdbcPersistentProperty.java index 1bcc1be8..841d64c6 100644 --- a/src/main/java/org/springframework/data/jdbc/mapping/model/JdbcPersistentProperty.java +++ b/src/main/java/org/springframework/data/jdbc/mapping/model/JdbcPersistentProperty.java @@ -23,6 +23,8 @@ import org.springframework.data.mapping.model.AnnotationBasedPersistentProperty; import org.springframework.data.mapping.model.SimpleTypeHolder; /** + * meta data about a property to be used by repository implementations. + * * @author Jens Schauder */ public class JdbcPersistentProperty extends AnnotationBasedPersistentProperty { @@ -43,4 +45,8 @@ public class JdbcPersistentProperty extends AnnotationBasedPersistentProperty createAssociation() { return null; } + + public String getColumnName() { + return getName(); + } } diff --git a/src/main/java/org/springframework/data/jdbc/repository/EntityRowMapper.java b/src/main/java/org/springframework/data/jdbc/repository/EntityRowMapper.java index 70450670..4bde019f 100644 --- a/src/main/java/org/springframework/data/jdbc/repository/EntityRowMapper.java +++ b/src/main/java/org/springframework/data/jdbc/repository/EntityRowMapper.java @@ -15,49 +15,62 @@ */ package org.springframework.data.jdbc.repository; -import java.lang.reflect.InvocationTargetException; import java.sql.ResultSet; import java.sql.SQLException; -import org.springframework.data.mapping.PersistentEntity; +import org.springframework.data.convert.ClassGeneratingEntityInstantiator; +import org.springframework.data.convert.EntityInstantiator; +import org.springframework.data.jdbc.mapping.model.JdbcPersistentEntity; +import org.springframework.data.jdbc.mapping.model.JdbcPersistentProperty; import org.springframework.data.mapping.PersistentProperty; +import org.springframework.data.mapping.PreferredConstructor; import org.springframework.data.mapping.PropertyHandler; +import org.springframework.data.mapping.model.MappingException; +import org.springframework.data.mapping.model.ParameterValueProvider; /** + * maps a ResultSet to an entity of type {@code T} + * * @author Jens Schauder */ class EntityRowMapper implements org.springframework.jdbc.core.RowMapper { - private final PersistentEntity entity; + private final JdbcPersistentEntity entity; - EntityRowMapper(PersistentEntity entity) { + private final EntityInstantiator instantiator = new ClassGeneratingEntityInstantiator(); + + EntityRowMapper(JdbcPersistentEntity entity) { this.entity = entity; } @Override public T mapRow(ResultSet rs, int rowNum) throws SQLException { - try { + T t = createInstance(rs); - T t = createInstance(); + entity.doWithProperties((PropertyHandler) property -> { + setProperty(rs, t, property); + }); - entity.doWithProperties((PropertyHandler) property -> { - setProperty(rs, t, property); - }); - - return t; - } catch (Exception e) { - throw new RuntimeException(String.format("Could not instantiate %s", entity.getType())); - } + return t; } - private T createInstance() throws InstantiationException, IllegalAccessException, InvocationTargetException { - return (T) entity.getPersistenceConstructor().getConstructor().newInstance(); + private T createInstance(ResultSet rs) { + return instantiator.createInstance(entity, new ParameterValueProvider() { + @Override + public T getParameterValue(PreferredConstructor.Parameter parameter) { + try { + return (T) rs.getObject(parameter.getName()); + } catch (SQLException e) { + throw new MappingException(String.format("Couldn't read column %s from ResultSet.", parameter.getName())); + } + } + }); } private void setProperty(ResultSet rs, T t, PersistentProperty property) { try { - property.getSetter().invoke(t, rs.getObject(property.getName())); + entity.getPropertyAccessor(t).setProperty(property, rs.getObject(property.getName())); } catch (Exception e) { throw new RuntimeException(String.format("Couldn't set property %s.", property.getName()), e); } diff --git a/src/main/java/org/springframework/data/jdbc/repository/SimpleJdbcRepository.java b/src/main/java/org/springframework/data/jdbc/repository/SimpleJdbcRepository.java index 8ae8d4b7..556e2a1e 100644 --- a/src/main/java/org/springframework/data/jdbc/repository/SimpleJdbcRepository.java +++ b/src/main/java/org/springframework/data/jdbc/repository/SimpleJdbcRepository.java @@ -16,14 +16,13 @@ package org.springframework.data.jdbc.repository; import java.io.Serializable; -import java.util.ArrayList; import java.util.HashMap; -import java.util.List; import java.util.Map; import java.util.stream.Collectors; +import java.util.stream.StreamSupport; import javax.sql.DataSource; import org.springframework.data.jdbc.mapping.model.JdbcPersistentEntity; -import org.springframework.data.mapping.PersistentProperty; +import org.springframework.data.jdbc.mapping.model.JdbcPersistentProperty; import org.springframework.data.mapping.PropertyHandler; import org.springframework.data.repository.CrudRepository; import org.springframework.jdbc.core.namedparam.MapSqlParameterSource; @@ -37,117 +36,116 @@ public class SimpleJdbcRepository implements CrudRep private final JdbcPersistentEntity entity; private final NamedParameterJdbcOperations template; + private final SqlGenerator sql; - private final String findOneSql; - private final String insertSql; + private final EntityRowMapper entityRowMapper; public SimpleJdbcRepository(JdbcPersistentEntity entity, DataSource dataSource) { this.entity = entity; this.template = new NamedParameterJdbcTemplate(dataSource); - findOneSql = createFindOneSelectSql(); - insertSql = createInsertSql(); + entityRowMapper = new EntityRowMapper(entity); + sql = new SqlGenerator(entity); } @Override public S save(S entity) { - template.update(insertSql, getPropertyMap(entity)); + template.update(sql.getInsert(), getPropertyMap(entity)); return entity; } @Override public Iterable save(Iterable entities) { - return null; + + Map[] batchValues = StreamSupport + .stream(entities.spliterator(), false) + .map(i -> getPropertyMap(i)) + .toArray(size -> new Map[size]); + + template.batchUpdate(sql.getInsert(), batchValues); + + return entities; } @Override public T findOne(ID id) { return template.queryForObject( - findOneSql, + sql.getFindOne(), new MapSqlParameterSource("id", id), - new EntityRowMapper(entity) + entityRowMapper ); } @Override public boolean exists(ID id) { - return false; + + return template.queryForObject( + sql.getExists(), + new MapSqlParameterSource("id", id), + Boolean.class + ); } @Override public Iterable findAll() { - return null; + return template.query(sql.getFindAll(), entityRowMapper); } @Override public Iterable findAll(Iterable ids) { - return null; + return template.query(sql.getFindAllInList(), new MapSqlParameterSource("ids", ids), entityRowMapper); } @Override public long count() { - return 0; + return template.getJdbcOperations().queryForObject(sql.getCount(), Long.class); } @Override public void delete(ID id) { - + template.update(sql.getDeleteById(), new MapSqlParameterSource("id", id)); } @Override - public void delete(T entity) { + public void delete(T instance) { + template.update( + sql.getDeleteById(), + new MapSqlParameterSource("id", + entity.getIdValue(instance))); } @Override public void delete(Iterable entities) { + template.update( + sql.getDeleteByList(), + new MapSqlParameterSource("ids", + StreamSupport + .stream(entities.spliterator(), false) + .map(entity::getIdValue) + .collect(Collectors.toList()) + ) + ); } @Override public void deleteAll() { - + template.getJdbcOperations().update(sql.getDeleteAll()); } - private String createFindOneSelectSql() { - - String tableName = entity.getType().getSimpleName(); - String idColumn = entity.getIdProperty().getName(); - - return String.format("select * from %s where %s = :id", tableName, idColumn); - } - - private String createInsertSql() { - - List propertyNames = new ArrayList<>(); - entity.doWithProperties((PropertyHandler) persistentProperty -> propertyNames.add(persistentProperty.getName())); - - String insertTemplate = "insert into %s (%s) values (%s)"; - - String tableName = entity.getType().getSimpleName(); - - String tableColumns = propertyNames.stream().collect(Collectors.joining(", ")); - String parameterNames = propertyNames.stream().collect(Collectors.joining(", :", ":", "")); - - return String.format(insertTemplate, tableName, tableColumns, parameterNames); - } - - private Map getPropertyMap(final S entity) { + private Map getPropertyMap(final S instance) { Map parameters = new HashMap<>(); - this.entity.doWithProperties(new PropertyHandler() { + this.entity.doWithProperties(new PropertyHandler() { @Override - public void doWithPersistentProperty(PersistentProperty persistentProperty) { - try { - parameters.put(persistentProperty.getName(), persistentProperty.getGetter().invoke(entity)); - } catch (Exception e) { - throw new RuntimeException(String.format("Couldn't get value of property %s", persistentProperty.getName())); - } + public void doWithPersistentProperty(JdbcPersistentProperty persistentProperty) { + parameters.put(persistentProperty.getColumnName(), entity.getPropertyAccessor(instance).getProperty(persistentProperty)); } }); diff --git a/src/main/java/org/springframework/data/jdbc/repository/SqlGenerator.java b/src/main/java/org/springframework/data/jdbc/repository/SqlGenerator.java new file mode 100644 index 00000000..b561029b --- /dev/null +++ b/src/main/java/org/springframework/data/jdbc/repository/SqlGenerator.java @@ -0,0 +1,138 @@ +/* + * Copyright 2017 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.data.jdbc.repository; + +import java.util.ArrayList; +import java.util.List; +import java.util.stream.Collectors; +import org.springframework.data.jdbc.mapping.model.JdbcPersistentEntity; +import org.springframework.data.mapping.PropertyHandler; + +/** + * @author Jens Schauder + */ +class SqlGenerator { + + private final String findOneSql; + private final String findAllSql; + private final String findAllInListSql; + + private final String existsSql; + private final String countSql; + + private final String insertSql; + private final String deleteByIdSql; + private final String deleteAllSql; + private final String deleteByListSql; + + SqlGenerator(JdbcPersistentEntity entity) { + + findOneSql = createFindOneSelectSql(entity); + findAllSql = createFindAllSql(entity); + findAllInListSql = createFindAllInListSql(entity); + + existsSql = createExistsSql(entity); + countSql = createCountSql(entity); + + insertSql = createInsertSql(entity); + + deleteByIdSql = createDeleteSql(entity); + deleteAllSql = createDeleteAllSql(entity); + deleteByListSql = createDeleteByListSql(entity); + } + + String getFindAllInList() { + return findAllInListSql; + } + + String getFindAll() { + return findAllSql; + } + + String getExists() { + return existsSql; + } + + String getFindOne() { + return findOneSql; + } + + String getInsert() { + return insertSql; + } + + String getCount() { + return countSql; + } + + String getDeleteById() { + return deleteByIdSql; + } + + String getDeleteAll() { + return deleteAllSql; + } + + String getDeleteByList() { + return deleteByListSql; + } + private String createFindOneSelectSql(JdbcPersistentEntity entity) { + return String.format("select * from %s where %s = :id", entity.getTableName(), entity.getIdColumn()); + } + + private String createFindAllSql(JdbcPersistentEntity entity) { + return String.format("select * from %s", entity.getTableName()); + } + + private String createFindAllInListSql(JdbcPersistentEntity entity) { + return String.format(String.format("select * from %s where %s in (:ids)", entity.getTableName(), entity.getIdColumn()), entity.getTableName()); + } + + private String createExistsSql(JdbcPersistentEntity entity) { + return String.format("select count(*) from %s where %s = :id", entity.getTableName(), entity.getIdColumn()); + } + + private String createCountSql(JdbcPersistentEntity entity) { + return String.format("select count(*) from %s", entity.getTableName(), entity.getIdColumn()); + } + + private String createInsertSql(JdbcPersistentEntity entity) { + + List propertyNames = new ArrayList<>(); + entity.doWithProperties((PropertyHandler) persistentProperty -> propertyNames.add(persistentProperty.getName())); + + String insertTemplate = "insert into %s (%s) values (%s)"; + + String tableName = entity.getType().getSimpleName(); + + String tableColumns = propertyNames.stream().collect(Collectors.joining(", ")); + String parameterNames = propertyNames.stream().collect(Collectors.joining(", :", ":", "")); + + return String.format(insertTemplate, tableName, tableColumns, parameterNames); + } + + private String createDeleteSql(JdbcPersistentEntity entity) { + return String.format("delete from %s where %s = :id", entity.getTableName(), entity.getIdColumn()); + } + + private String createDeleteAllSql(JdbcPersistentEntity entity) { + return String.format("delete from %s", entity.getTableName()); + } + + private String createDeleteByListSql(JdbcPersistentEntity entity) { + return String.format("delete from %s where id in (:ids)", entity.getTableName()); + } +} diff --git a/src/test/java/org/springframework/data/jdbc/repository/JdbcRepositoryIntegrationTests.java b/src/test/java/org/springframework/data/jdbc/repository/JdbcRepositoryIntegrationTests.java index b1fd9045..8af0faf1 100644 --- a/src/test/java/org/springframework/data/jdbc/repository/JdbcRepositoryIntegrationTests.java +++ b/src/test/java/org/springframework/data/jdbc/repository/JdbcRepositoryIntegrationTests.java @@ -15,10 +15,13 @@ */ package org.springframework.data.jdbc.repository; +import static java.util.Arrays.*; +import static org.assertj.core.api.Assertions.assertThat; import static org.junit.Assert.*; import java.sql.SQLException; import org.junit.After; +import org.junit.Assert; import org.junit.Test; import org.springframework.data.annotation.Id; import org.springframework.data.jdbc.repository.support.JdbcRepositoryFactory; @@ -50,7 +53,7 @@ public class JdbcRepositoryIntegrationTests { private final DummyEntityRepository repository = createRepository(db); - private DummyEntity entity = createDummyEntity(); + private DummyEntity entity = createDummyEntity(23L); @After public void after() { @@ -59,7 +62,7 @@ public class JdbcRepositoryIntegrationTests { @Test - public void canSaveAnEntity() throws SQLException { + public void canSaveAnEntity() { entity = repository.save(entity); @@ -74,7 +77,7 @@ public class JdbcRepositoryIntegrationTests { } @Test - public void canSaveAndLoadAnEntity() throws SQLException { + public void canSaveAndLoadAnEntity() { entity = repository.save(entity); @@ -88,14 +91,120 @@ public class JdbcRepositoryIntegrationTests { reloadedEntity.getName()); } + @Test + public void saveMany() { + + DummyEntity other = createDummyEntity(24L); + + repository.save(asList(entity, other)); + + assertThat(repository.findAll()).extracting(DummyEntity::getId).containsExactlyInAnyOrder(23L, 24L); + } + + @Test + public void existsReturnsTrueIffEntityExists() { + + entity = repository.save(entity); + + assertTrue(repository.exists(entity.getId())); + assertFalse(repository.exists(entity.getId() + 1)); + } + + @Test + public void findAllFindsAllEntities() { + + DummyEntity other = createDummyEntity(24L); + + other = repository.save(other); + entity = repository.save(entity); + + Iterable all = repository.findAll(); + + assertThat(all).extracting("id").containsExactlyInAnyOrder(entity.getId(), other.getId()); + } + + @Test + public void findAllFindsAllSpecifiedEntities() { + + repository.save(createDummyEntity(24L)); + DummyEntity other = repository.save(createDummyEntity(25L)); + entity = repository.save(entity); + + Iterable all = repository.findAll(asList(entity.getId(), other.getId())); + + assertThat(all).extracting("id").containsExactlyInAnyOrder(entity.getId(), other.getId()); + } + + @Test + public void count() { + + repository.save(createDummyEntity(24L)); + repository.save(createDummyEntity(25L)); + repository.save(entity); + + assertThat(repository.count()).isEqualTo(3L); + } + + @Test + public void deleteById() { + + repository.save(createDummyEntity(24L)); + repository.save(createDummyEntity(25L)); + repository.save(entity); + + repository.delete(24L); + + assertThat(repository.findAll()).extracting(DummyEntity::getId).containsExactlyInAnyOrder(23L, 25L); + } + + @Test + public void deleteByEntity() { + + repository.save(createDummyEntity(24L)); + repository.save(createDummyEntity(25L)); + repository.save(entity); + + repository.delete(entity); + + assertThat(repository.findAll()).extracting(DummyEntity::getId).containsExactlyInAnyOrder(24L, 25L); + } + + + @Test + public void deleteByList() { + + repository.save(entity); + repository.save(createDummyEntity(24L)); + DummyEntity other = repository.save(createDummyEntity(25L)); + + repository.delete(asList(entity, other)); + + assertThat(repository.findAll()).extracting(DummyEntity::getId).containsExactlyInAnyOrder(24L); + } + + @Test + public void deleteAll() { + + repository.save(entity); + repository.save(createDummyEntity(24L)); + repository.save(createDummyEntity(25L)); + + repository.deleteAll(); + + assertThat(repository.findAll()).isEmpty(); + } + + + private static DummyEntityRepository createRepository(EmbeddedDatabase db) { return new JdbcRepositoryFactory(db).getRepository(DummyEntityRepository.class); } - private static DummyEntity createDummyEntity() { + private static DummyEntity createDummyEntity(long id) { + DummyEntity entity = new DummyEntity(); - entity.setId(23L); + entity.setId(id); entity.setName("Entity Name"); return entity; } @@ -104,8 +213,9 @@ public class JdbcRepositoryIntegrationTests { } + // needs to be public in order for the Hamcrest property matcher to work. @Data - private static class DummyEntity { + public static class DummyEntity { @Id Long id;