diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/config/AbstractCassandraConfiguration.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/config/AbstractCassandraConfiguration.java index c43802432..f66e534bc 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/config/AbstractCassandraConfiguration.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/config/AbstractCassandraConfiguration.java @@ -15,6 +15,7 @@ */ package org.springframework.data.cassandra.config; +import java.util.Arrays; import java.util.Collections; import java.util.Optional; import java.util.Set; @@ -227,7 +228,12 @@ public abstract class AbstractCassandraConfiguration extends AbstractSessionConf * @since 2.0 */ protected Set> getInitialEntitySet() throws ClassNotFoundException { - return CassandraEntityClassScanner.scan(getEntityBasePackages()); + + CassandraEntityClassScanner scanner = new CassandraEntityClassScanner(); + scanner.setBeanClassLoader(this.beanClassLoader); + scanner.setEntityBasePackages(Arrays.asList(getEntityBasePackages())); + + return scanner.scanForEntityClasses(); } /** diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/config/CassandraEntityClassScanner.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/config/CassandraEntityClassScanner.java index e132327af..a9409410b 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/config/CassandraEntityClassScanner.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/config/CassandraEntityClassScanner.java @@ -21,13 +21,11 @@ import java.util.Collection; import java.util.Collections; import java.util.HashSet; import java.util.Set; +import java.util.stream.Collectors; -import org.springframework.beans.factory.config.BeanDefinition; -import org.springframework.context.annotation.ClassPathScanningCandidateComponentProvider; -import org.springframework.core.type.filter.AnnotationTypeFilter; -import org.springframework.data.annotation.Persistent; import org.springframework.data.cassandra.core.mapping.PrimaryKeyClass; import org.springframework.data.cassandra.core.mapping.Table; +import org.springframework.data.util.TypeScanner; import org.springframework.lang.Nullable; import org.springframework.util.ClassUtils; import org.springframework.util.ObjectUtils; @@ -145,7 +143,12 @@ public class CassandraEntityClassScanner { * @return base package names used for the entity scan. */ public Set getEntityBasePackages() { - return Collections.unmodifiableSet(entityBasePackages); + + if (ObjectUtils.isEmpty(entityBasePackageClasses)) { + return Collections.unmodifiableSet(entityBasePackages); + } + + return entityBasePackageClasses.stream().map(ClassUtils::getPackageName).collect(Collectors.toSet()); } /** @@ -176,57 +179,30 @@ public class CassandraEntityClassScanner { /** * Set the bean {@link ClassLoader} to load class candidates discovered by the class path scan. * - * @param beanClassLoader must not be {@literal null}. + * @param beanClassLoader */ - public void setBeanClassLoader(ClassLoader beanClassLoader) { + public void setBeanClassLoader(@Nullable ClassLoader beanClassLoader) { this.beanClassLoader = beanClassLoader; } /** - * Scans the mapping base package for entity classes annotated with {@link Table} or {@link Persistent}. + * Scans the mapping base package for entity classes annotated with {@link Table} or {@link PrimaryKeyClass}. * * @see #getEntityBasePackages() - * @return {@code Set>} representing the annotated entity classes found. - * @throws ClassNotFoundException if a discovered class could not be loaded via. + * @see #getEntityAnnotations() + * @return {@code Set} representing the annotated entity classes found. */ - public Set> scanForEntityClasses() throws ClassNotFoundException { + public Set> scanForEntityClasses() { - Set> classes = new HashSet<>(); + TypeScanner scanner; - for (String basePackage : getEntityBasePackages()) { - classes.addAll(scanBasePackageForEntities(basePackage)); + if (this.beanClassLoader != null) { + scanner = TypeScanner.typeScanner(this.beanClassLoader); + } else { + scanner = TypeScanner.typeScanner(ClassUtils.getDefaultClassLoader()); } - for (Class basePackageClass : getEntityBasePackageClasses()) { - classes.addAll(scanBasePackageForEntities(basePackageClass.getPackage().getName())); - } - - return classes; - } - - protected Set> scanBasePackageForEntities(String basePackage) throws ClassNotFoundException { - - HashSet> classes = new HashSet<>(); - - if (ObjectUtils.isEmpty(basePackage)) { - return classes; - } - - ClassPathScanningCandidateComponentProvider componentProvider = new ClassPathScanningCandidateComponentProvider( - false); - - for (Class annotation : getEntityAnnotations()) { - componentProvider.addIncludeFilter(new AnnotationTypeFilter(annotation)); - } - - for (BeanDefinition candidate : componentProvider.findCandidateComponents(basePackage)) { - - if (candidate.getBeanClassName() != null) { - classes.add(ClassUtils.forName(candidate.getBeanClassName(), beanClassLoader)); - } - } - - return classes; + return scanner.forTypesAnnotatedWith(getEntityAnnotations()).scanPackages(getEntityBasePackages()).collectAsSet(); } /** diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/config/CassandraMappingContextParser.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/config/CassandraMappingContextParser.java index 51f9f0c3c..ae2d5d5d7 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/config/CassandraMappingContextParser.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/config/CassandraMappingContextParser.java @@ -61,18 +61,21 @@ class CassandraMappingContextParser extends AbstractSingleBeanDefinitionParser { @Override protected void doParse(Element element, ParserContext parserContext, BeanDefinitionBuilder builder) { - parseMapping(element, builder); + parseMapping(element, builder, parserContext.getReaderContext().getBeanClassLoader()); builder.getRawBeanDefinition().setSource(element); } - private void parseMapping(Element element, BeanDefinitionBuilder builder) { + private void parseMapping(Element element, BeanDefinitionBuilder builder, ClassLoader classLoader) { String packages = element.getAttribute("entity-base-packages"); if (StringUtils.hasText(packages)) { try { - Set> entityClasses = CassandraEntityClassScanner - .scan(StringUtils.commaDelimitedListToStringArray(packages)); + CassandraEntityClassScanner scanner = new CassandraEntityClassScanner(); + scanner.setBeanClassLoader(classLoader); + scanner.setEntityBasePackages(StringUtils.commaDelimitedListToSet(packages)); + + Set> entityClasses = scanner.scanForEntityClasses(); builder.addPropertyValue("initialEntitySet", entityClasses); } catch (Exception x) {