View Javadoc
1   package com.reallifedeveloper.tools.test.database;
2   
3   import java.io.Serializable;
4   import java.lang.reflect.Constructor;
5   import java.lang.reflect.Field;
6   import java.lang.reflect.ParameterizedType;
7   import java.lang.reflect.Type;
8   import java.math.BigDecimal;
9   import java.math.BigInteger;
10  import java.time.Instant;
11  import java.time.LocalDate;
12  import java.time.LocalDateTime;
13  import java.time.ZonedDateTime;
14  import java.util.ArrayList;
15  import java.util.Arrays;
16  import java.util.Collection;
17  import java.util.Date;
18  import java.util.HashMap;
19  import java.util.HashSet;
20  import java.util.List;
21  import java.util.Map;
22  import java.util.Optional;
23  import java.util.Set;
24  import java.util.UUID;
25  
26  import org.checkerframework.checker.nullness.qual.Nullable;
27  import org.slf4j.Logger;
28  import org.slf4j.LoggerFactory;
29  import org.springframework.data.repository.CrudRepository;
30  
31  import edu.umd.cs.findbugs.annotations.SuppressFBWarnings;
32  import jakarta.persistence.Embeddable;
33  import jakarta.persistence.Entity;
34  import jakarta.persistence.Id;
35  import jakarta.persistence.JoinColumn;
36  import jakarta.persistence.JoinTable;
37  import jakarta.persistence.OneToMany;
38  import jakarta.persistence.OneToOne;
39  import lombok.Getter;
40  
41  import com.reallifedeveloper.tools.test.TestUtil;
42  
43  /**
44   * A helper class to write data into a {@link CrudRepository} from some data source, e.g., a CSV file, where each entity is represented by a
45   * {@link DbTableRow}.
46   * <p>
47   * This can be useful for inserting test data into a repository, irrespective of whether the repository connects to a real database or not.
48   * <p>
49   * TODO: The current implementation only has basic support for "to many" associations (there must be a &amp;JoinTable annotation on a field,
50   * with &amp;JoinColumn annotations), and for enums (an enum must be stored as a string).
51   *
52   * @author RealLifeDeveloper
53   */
54  @Getter
55  @SuppressWarnings("PMD")
56  @SuppressFBWarnings(value = { "CRLF_INJECTION_LOGS", "IMPROPER_UNICODE" })
57  public class CrudRepositoryWriter {
58  
59      private static final Logger LOG = LoggerFactory.getLogger(CrudRepositoryWriter.class);
60  
61      private final EntityMap entityMap = new EntityMap();
62  
63      /**
64       * Creates a new entity based on data from a {@link DbTableRow} and writes it into a repository if appropriate.
65       * <p>
66       * This method may create entities that are not directly handled by the repository, in which case they are assumed to be related to some
67       * entity in the repository.
68       *
69       * @param <T>                  the type of entities in the repository
70       * @param <E>                  the type of entity being created
71       * @param <ID>                 the type of the primary key of the entities in the repository
72       * @param tableRow             the {@code TableRow} with the data to insert into the fields of the newly created entity
73       * @param repositoryEntityType the class object representing {@code T}, i.e., the type ofrepository entities
74       * @param entityType           the class object representing {@code E}, i.e., the type of entity being created, or {@code null}
75       * @param repository           the repository in which to insert the newly created entity
76       * @param tableName            the name of the database table where the entity should be stored
77       * @return {@code true} if an entity was created, no matter if it was saved in the repository, {@code false} otherwise
78       * @throws ReflectiveOperationException if some reflection operation failed creating the entity or setting is fields
79       */
80      public <T, E, ID extends Serializable> boolean writeEntity(DbTableRow tableRow, Class<T> repositoryEntityType,
81              @Nullable Class<E> entityType, CrudRepository<T, ID> repository, String tableName) throws ReflectiveOperationException {
82          if (entityType == null) {
83              return false;
84          }
85          if (entityType.getAnnotation(Embeddable.class) != null) {
86              return writeEmbeddable(tableRow, repositoryEntityType, entityType, repository, tableName);
87          }
88          if (entityType.getAnnotation(Entity.class) == null || !JpaUtil.getTableName(entityType).equalsIgnoreCase(tableName)) {
89              return false;
90          }
91          E entity = createEntity(entityType);
92          for (DbTableField column : tableRow.columns()) {
93              String fieldName = JpaUtil.getFieldName(column.name(), entityType);
94              setField(entity, fieldName, column.value());
95          }
96          entityMap.addEntity(entity);
97          saveToRepository(entity, repository, repositoryEntityType);
98          return true;
99      }
100 
101     private <T, ID extends Serializable> void saveToRepository(Object entity, CrudRepository<T, ID> repository,
102             Class<T> repositoryEntityType) {
103         if (repositoryEntityType.isAssignableFrom(entity.getClass())) {
104             T entityToSave = repositoryEntityType.cast(entity);
105             LOG.debug("Saving entity in repository: entity={}", entity);
106             repository.save(entityToSave);
107         }
108     }
109 
110     @SuppressWarnings("UnusedVariable")
111     private <T, E, ID extends Serializable> boolean writeEmbeddable(DbTableRow tableRow, Class<T> repositoryEntityType,
112             @Nullable Class<E> entityType, CrudRepository<T, ID> repository, String tableName) {
113         LOG.debug("Saving embeddable {}", tableRow);
114         throw new UnsupportedOperationException("writeEmbeddable not yet implemented");
115     }
116 
117     /**
118      * Connects entities based on data in a join table.
119      *
120      * @param tableRow       the {@code TableRow} with the data for the join table
121      * @param joinTtableName the name of the join table to use to connect entities
122      */
123     public void addEntitiesFromJoinTable(DbTableRow tableRow, String joinTtableName) {
124         entityMap.joinTableField(joinTtableName).ifPresent(joinTableField -> {
125             joinTableField.setAccessible(true);
126             ParameterizedType parameterizedType = (ParameterizedType) joinTableField.getGenericType();
127             Class<?> targetType = (Class<?>) parameterizedType.getActualTypeArguments()[0];
128             JoinTable joinTable = joinTableField.getAnnotation(JoinTable.class);
129             assert joinTable != null : "JoinTable annotation should be present when the joinTableField method returns a non-empty value";
130             for (JoinColumn joinColumn : joinTable.joinColumns()) {
131                 for (JoinColumn inverseJoinColumn : joinTable.inverseJoinColumns()) {
132                     addEntityFromJoinTable(tableRow, joinTableField, targetType, joinColumn, inverseJoinColumn);
133                 }
134             }
135         });
136     }
137 
138     /**
139      * Goes through all entities that have been saved, trying to fix missing associations.
140      *
141      * @param <T>                  the type of entities in the repository
142      * @param <ID>                 the type of the primary key of the entities in the repository
143      *
144      * @param repository           the repository in which to save the entities that may have been updated
145      * @param repositoryEntityType the class object representing {@code T}, i.e., the type ofrepository entities
146      *
147      * @throws ReflectiveOperationException if something went wrong using reflection to analyze the entities
148      */
149     public <T, ID extends Serializable> void fillReferencesBetweenEntities(CrudRepository<T, ID> repository, Class<T> repositoryEntityType)
150             throws ReflectiveOperationException {
151         for (Object entity : entityMap.entities()) {
152             LOG.trace("fillReferencesBetweenEntities: Examining entity {}", entity);
153             for (Field field : entity.getClass().getDeclaredFields()) {
154                 field.setAccessible(true);
155                 OneToOne oneToOne = field.getAnnotation(OneToOne.class);
156                 if (oneToOne != null) {
157                     handleOneToOne(entity, field, oneToOne);
158                 }
159                 OneToMany oneToMany = field.getAnnotation(OneToMany.class);
160                 if (oneToMany != null) {
161                     handleOneToMany(entity, field, oneToMany);
162                 }
163             }
164             saveToRepository(entity, repository, repositoryEntityType);
165         }
166     }
167 
168     private void handleOneToOne(Object entity, Field field, OneToOne oneToOne) throws IllegalAccessException, NoSuchFieldException {
169         if (field.get(entity) != null) {
170             LOG.debug("handleOneToOne: Field already set, ignoring it: field={}, entity={}", field, entity);
171             return;
172         }
173         if (field.getAnnotation(JoinColumn.class) != null) {
174             LOG.trace("handleOneToOne: Ignoring owner side, it should be set as normal field: field={}, entity={}", field, entity);
175         } else if (oneToOne.mappedBy() != null && !oneToOne.mappedBy().isEmpty()) {
176             LOG.trace("handleOneToOne: Inverse side: field={}, entity={}", field, entity);
177             String mappedBy = oneToOne.mappedBy();
178             Object id = JpaUtil.getIdValue(entity);
179             List<?> entitiesToMap = findEntitiesByClassAndField(field.getType(), mappedBy, id);
180             if (entitiesToMap.isEmpty()) {
181                 LOG.trace("handleOneToOne: Inverse side found no candidate for OneToOne mapping: field={}, entity={}", field, entity);
182                 return;
183             } else if (entitiesToMap.size() > 1) {
184                 throw new IllegalStateException("Found multiple candidates for OneToOne mapping: field=" + field + ", entity=" + entity);
185             }
186             Object value = entitiesToMap.get(0);
187             LOG.debug("Setting OneToOne field {} to {}", JpaUtil.fieldNameForLogging(entity, field), value);
188             if (value == null) {
189                 LOG.warn("handleOneToOne: Not setting OneToOneField to null: field={}, entity={}", field, entity);
190                 return;
191             }
192             field.set(entity, value);
193         } else {
194             throw new IllegalStateException(
195                     "OneToOne field " + JpaUtil.fieldNameForLogging(entity, field) + " has no mappedBy and no JoinColumn annotation");
196         }
197     }
198 
199     private void handleOneToMany(Object entity, Field field, OneToMany oneToMany) throws IllegalAccessException, NoSuchFieldException {
200         Class<?> collectionType = field.getType();
201         if (Collection.class.isAssignableFrom(collectionType)) {
202             saveEntityInCollection(entity, field, oneToMany);
203         } else if (Map.class.isAssignableFrom(collectionType)) {
204             saveEntityInMap(entity, field, oneToMany);
205         }
206     }
207 
208     private void saveEntityInCollection(Object entity, Field field, OneToMany oneToMany)
209             throws IllegalAccessException, NoSuchFieldException {
210         LOG.trace("saveEntityInCollection: entity={}, field={}, oneToMany={}", entity, field, oneToMany);
211         LOG.trace("Not yet implemented");
212     }
213 
214     private void saveEntityInMap(Object entity, Field field, OneToMany oneToMany) throws IllegalAccessException, NoSuchFieldException {
215         LOG.trace("saveEntityInMap: entity={}, field={}, oneToMany={}", entity, field, oneToMany);
216         ParameterizedType parameterizedType = (ParameterizedType) field.getGenericType();
217         Type[] targetTypes = parameterizedType.getActualTypeArguments();
218         Class<?> targetClass = getClass(targetTypes[1].getTypeName());
219         String mappedBy = oneToMany.mappedBy();
220         if (mappedBy == null || mappedBy.isEmpty()) {
221             JoinColumn joinColumn = field.getAnnotation(JoinColumn.class);
222             if (joinColumn != null) {
223                 mappedBy = JpaUtil.getFieldName(joinColumn.name(), targetClass);
224             }
225         }
226         if (mappedBy == null || mappedBy.isEmpty()) {
227             throw new IllegalStateException(
228                     "OneToMany field " + JpaUtil.fieldNameForLogging(entity, field) + " has no mappedBy and no JoinColumn annotation");
229         }
230         Object id = JpaUtil.getIdValue(entity);
231         List<?> entitiesToMap = findEntitiesByClassAndField(targetClass, mappedBy, id);
232         JpaUtil.addEntitiesToMapField(field, entity, entitiesToMap);
233     }
234 
235     private static Class<?> getClass(String className) {
236         try {
237             return Class.forName(className);
238         } catch (ClassNotFoundException e) {
239             throw new IllegalStateException("Class " + className + " not found", e);
240         }
241     }
242 
243     private <T> List<T> findEntitiesByClassAndField(Class<T> entityClass, String fieldName, Object value)
244             throws IllegalAccessException, NoSuchFieldException {
245         LOG.trace("Finding entities by class={}, field={} and value={}", entityClass, fieldName, value);
246         List<T> foundEntities = new ArrayList<>();
247         for (T entity : entityMap.entitiesOfType(entityClass)) {
248             Field field = entity.getClass().getDeclaredField(fieldName);
249             field.setAccessible(true);
250             // LOG.debug("{}.{}={}", entityClass.getName(), fieldName, field.get(entity));
251             Object fieldValue = field.get(entity);
252             if (fieldValue == null) {
253                 continue;
254             }
255             if (fieldValue.equals(value)) {
256                 foundEntities.add(entity);
257             } else if (fieldValue.getClass().getAnnotation(Entity.class) != null && JpaUtil.getIdValue(fieldValue).equals(value)) {
258                 foundEntities.add(entity);
259             }
260         }
261         return foundEntities;
262     }
263 
264     private void addEntityFromJoinTable(DbTableRow tableRow, Field joinTableField, Class<?> targetType, JoinColumn joinColumn,
265             JoinColumn inverseJoinColumn) {
266         String lhsPrimaryKey = null;
267         String rhsPrimaryKey = null;
268         for (DbTableField column : tableRow.columns()) {
269             if (column.name().equalsIgnoreCase(joinColumn.name())) {
270                 lhsPrimaryKey = column.value();
271             } else if (column.name().equalsIgnoreCase(inverseJoinColumn.name())) {
272                 rhsPrimaryKey = column.value();
273             }
274         }
275         if (lhsPrimaryKey == null || rhsPrimaryKey == null) {
276             throw new IllegalStateException("Failed to find join table: missing attribute in DBUnit XML file: '" + joinColumn.name()
277                     + "' or '" + inverseJoinColumn.name() + "'");
278         }
279         Object lhs = entityMap.findEntity(lhsPrimaryKey, joinTableField.getDeclaringClass());
280         Object rhs = entityMap.findEntity(rhsPrimaryKey, targetType);
281         JpaUtil.addObjectToCollectionField(joinTableField, lhs, rhs);
282     }
283 
284     private <T> T createEntity(Class<T> entityType) throws ReflectiveOperationException {
285         Constructor<T> constructor = entityType.getDeclaredConstructor();
286         constructor.setAccessible(true);
287         return constructor.newInstance();
288     }
289 
290     private <T> void setField(T entity, String fieldName, String attributeValue) throws ReflectiveOperationException {
291         Field field = JpaUtil.getField(entity, fieldName);
292         field.setAccessible(true);
293         Object fieldValue = createObjectFromString(attributeValue, field, JpaUtil.getPrimaryKeyType(entity.getClass()));
294         LOG.trace("Setting field {} to {}", fieldName, fieldValue);
295         field.set(entity, fieldValue);
296         if (fieldValue != null && fieldValue.getClass().getAnnotation(Entity.class) != null) {
297             potentiallyAddValueToCollection(fieldValue, fieldName, entity);
298         }
299     }
300 
301     private <T> void potentiallyAddValueToCollection(Object entity, String fieldName, T value) {
302         for (Field field : entity.getClass().getDeclaredFields()) {
303             field.setAccessible(true);
304             OneToMany oneToMany = field.getAnnotation(OneToMany.class);
305             if (oneToMany == null) {
306                 continue;
307             }
308             if (oneToMany.mappedBy().equals(fieldName)) {
309                 JpaUtil.addObjectToCollectionField(field, entity, value);
310             }
311         }
312     }
313 
314     private @Nullable Object createObjectFromString(String s, Field field, Class<?> primaryKeyType) {
315         Class<?> type;
316         if (field.getAnnotation(Id.class) != null) {
317             type = primaryKeyType;
318         } else {
319             type = field.getType();
320         }
321         return entityMap.createObjectFromString(s, type);
322     }
323 
324     /**
325      * Represents one row of data from the database.
326      *
327      * @author RealLifeDeveloper
328      */
329     public interface DbTableRow {
330         /**
331          * Gives the fields of this row.
332          *
333          * @return the fields
334          */
335         List<DbTableField> columns();
336     }
337 
338     /**
339      * Represents the value of a single field in the database.
340      *
341      * @param name  the name of the database column
342      * @param value the value of the field
343      */
344     public record DbTableField(String name, String value) {
345     }
346 
347     /**
348      * Keeps track of all entities handled by a {@code CrudRepositoryWriter}.
349      *
350      * @author RealLifeDeveloper
351      */
352     private static class EntityMap {
353 
354         private Map<PrimaryKey, Object> entities = new HashMap<>();
355         private final Set<Class<?>> entityClasses = new HashSet<>();
356 
357         void addEntity(Object entity) {
358             LOG.trace("Adding entity to internal map if necessary: entity={}", entity);
359             PrimaryKey primaryKey = PrimaryKey.fromEntity(entity);
360             if (!entities.containsKey(primaryKey)) {
361                 entities.put(primaryKey, entity);
362                 entityClasses.add(entity.getClass());
363             }
364         }
365 
366         Collection<Object> entities() {
367             return entities.values();
368         }
369 
370         Optional<Field> joinTableField(String tableName) {
371             for (Class<?> c : entityClasses) {
372                 for (Field field : c.getDeclaredFields()) {
373                     JoinTable joinTable = field.getAnnotation(JoinTable.class);
374                     if (joinTable != null && tableName.equalsIgnoreCase(joinTable.name())) {
375                         return Optional.of(field);
376                     }
377                 }
378             }
379             return Optional.empty();
380         }
381 
382         @SuppressWarnings("checkstyle:noReturnNull")
383         @Nullable
384         Object createObjectFromString(String s, Class<?> type) {
385             if (s == null || s.isEmpty()) {
386                 return null;
387             }
388             if (type == Byte.class) {
389                 return Byte.parseByte(s);
390             } else if (type == Short.class) {
391                 return Short.parseShort(s);
392             } else if (type == Integer.class) {
393                 return Integer.parseInt(s);
394             } else if (type == Long.class) {
395                 return Long.parseLong(s);
396             } else if (type == Float.class) {
397                 return Float.parseFloat(s);
398             } else if (type == Double.class) {
399                 return Double.parseDouble(s);
400             } else if (type == Boolean.class) {
401                 return Boolean.parseBoolean(s);
402             } else if (type == Character.class) {
403                 return s.charAt(0);
404             } else if (type == String.class) {
405                 return s;
406             } else if (type == Date.class) {
407                 return TestUtil.parseDate(s);
408             } else if (type == LocalDate.class) {
409                 return LocalDate.parse(s);
410             } else if (type == LocalDateTime.class) {
411                 return LocalDateTime.parse(s);
412             } else if (type == ZonedDateTime.class) {
413                 return ZonedDateTime.parse(s);
414             } else if (type == Instant.class) {
415                 return Instant.parse(s);
416             } else if (type == BigDecimal.class) {
417                 return new BigDecimal(s);
418             } else if (type == BigInteger.class) {
419                 return new BigInteger(s);
420             } else if (type == UUID.class) {
421                 return UUID.fromString(s);
422             } else if (type == List.class) {
423                 // Assuming Postgresql array syntax:
424                 return Arrays.asList(s.replaceAll("[{}]", "").split(","));
425             } else {
426                 return findEntity(s, type);
427             }
428         }
429 
430         @SuppressWarnings({ "rawtypes", "unchecked" })
431         Object findEntity(String strId, Class<?> entityType) {
432             if (entityType.isEnum()) {
433                 Class<? extends Enum> enumType = (Class<? extends Enum>) entityType;
434                 return Enum.valueOf(enumType, strId);
435             }
436             for (Object entity : entities()) {
437                 if (entity.getClass().equals(entityType)) {
438                     try {
439                         Object id = JpaUtil.getIdValue(entity);
440                         if (id != null && id.equals(createObjectFromString(strId, id.getClass()))) {
441                             return entity;
442                         }
443                     } catch (IllegalAccessException e) {
444                         throw new IllegalStateException(
445                                 "Unexpected problem looking up entity of " + entityType + " with primary key " + strId, e);
446                     }
447                 }
448             }
449             throw new IllegalArgumentException("Entity of " + entityType + " with primary key " + strId + " not found");
450         }
451 
452         @SuppressWarnings("unchecked")
453         private <T> List<T> entitiesOfType(Class<T> entityType) {
454             // LOG.debug("Getting entities of type {}", entityType.getName());
455             return (List<T>) entities().stream().filter(entity -> entity.getClass().equals(entityType)).toList();
456         }
457 
458         private record PrimaryKey(String entityClassName, Object id) {
459             static PrimaryKey fromEntity(Object entity) {
460                 try {
461                     return new PrimaryKey(entity.getClass().getName(), JpaUtil.getIdValue(entity));
462                 } catch (IllegalAccessException e) {
463                     throw new IllegalStateException("Unexpected problem creating PrimaryKey from entity: entity=" + entity, e);
464                 }
465             }
466         }
467     }
468 
469 }