diff --git a/src/main/java/org/springframework/hateoas/aot/AotUtils.java b/src/main/java/org/springframework/hateoas/aot/AotUtils.java index a05aa8b7..952c2da9 100644 --- a/src/main/java/org/springframework/hateoas/aot/AotUtils.java +++ b/src/main/java/org/springframework/hateoas/aot/AotUtils.java @@ -15,19 +15,30 @@ */ package org.springframework.hateoas.aot; +import java.io.IOException; +import java.util.Arrays; import java.util.HashSet; import java.util.List; import java.util.Optional; import java.util.Set; +import java.util.stream.Stream; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.aot.hint.MemberCategory; import org.springframework.aot.hint.ReflectionHints; +import org.springframework.aot.hint.TypeReference; +import org.springframework.beans.factory.config.BeanDefinition; +import org.springframework.context.annotation.ClassPathScanningCandidateComponentProvider; import org.springframework.core.ResolvableType; +import org.springframework.core.type.classreading.MetadataReader; +import org.springframework.core.type.classreading.MetadataReaderFactory; +import org.springframework.core.type.filter.AssignableTypeFilter; +import org.springframework.core.type.filter.TypeFilter; import org.springframework.hateoas.CollectionModel; import org.springframework.hateoas.EntityModel; import org.springframework.http.HttpEntity; +import org.springframework.util.ClassUtils; /** * Some helper classes to register types for reflection. @@ -115,4 +126,52 @@ class AotUtils { .flatMap(it -> extractGenerics(it, unresolved).stream()) .findFirst(); } + + public static FullTypeScanner getScanner(String packageName, TypeFilter... includeFilters) { + + var provider = new ClassPathScanningCandidateComponentProvider(false); + + if (includeFilters.length == 0) { + provider.addIncludeFilter(new AssignableTypeFilter(Object.class)); + } else { + Arrays.stream(includeFilters).forEach(provider::addIncludeFilter); + } + + provider.addExcludeFilter(new EnforcedPackageFilter(packageName)); + + return () -> provider.findCandidateComponents(packageName).stream() + .map(BeanDefinition::getBeanClassName) + .map(TypeReference::of); + } + + /** + * A {@link TypeFilter} to only match types outside the configured package. Usually used as exclude filter to + * limit scans to not find nested packages. + * + * @author Oliver Drotbohm + */ + private static class EnforcedPackageFilter implements TypeFilter { + + private final String referencePackage; + + public EnforcedPackageFilter(String referencePackage) { + this.referencePackage = referencePackage; + } + + /* + * (non-Javadoc) + * @see org.springframework.core.type.filter.TypeFilter#match(org.springframework.core.type.classreading.MetadataReader, org.springframework.core.type.classreading.MetadataReaderFactory) + */ + @Override + public boolean match(MetadataReader metadataReader, MetadataReaderFactory metadataReaderFactory) + throws IOException { + return !referencePackage + .equals(ClassUtils.getPackageName(metadataReader.getClassMetadata().getClassName())); + } + } + + static interface FullTypeScanner { + + abstract Stream findClasses(); + } } diff --git a/src/main/java/org/springframework/hateoas/aot/RepresentationModelRuntimeHints.java b/src/main/java/org/springframework/hateoas/aot/HateoasTypesRuntimeHints.java similarity index 60% rename from src/main/java/org/springframework/hateoas/aot/RepresentationModelRuntimeHints.java rename to src/main/java/org/springframework/hateoas/aot/HateoasTypesRuntimeHints.java index 41edb932..43430e85 100644 --- a/src/main/java/org/springframework/hateoas/aot/RepresentationModelRuntimeHints.java +++ b/src/main/java/org/springframework/hateoas/aot/HateoasTypesRuntimeHints.java @@ -15,16 +15,9 @@ */ package org.springframework.hateoas.aot; -import java.util.Arrays; -import java.util.List; -import java.util.stream.Stream; - import org.springframework.aot.hint.MemberCategory; import org.springframework.aot.hint.RuntimeHints; import org.springframework.aot.hint.RuntimeHintsRegistrar; -import org.springframework.hateoas.CollectionModel; -import org.springframework.hateoas.EntityModel; -import org.springframework.hateoas.PagedModel; import org.springframework.hateoas.RepresentationModel; /** @@ -32,13 +25,7 @@ import org.springframework.hateoas.RepresentationModel; * * @author Oliver Drotbohm */ -class RepresentationModelRuntimeHints implements RuntimeHintsRegistrar { - - private static final List> REPRESENTATION_MODELS = List.of(RepresentationModel.class, // - // EntityModel.class, // treated specially below - CollectionModel.class, // - PagedModel.class, - PagedModel.PageMetadata.class); +class HateoasTypesRuntimeHints implements RuntimeHintsRegistrar { /* * (non-Javadoc) @@ -48,11 +35,11 @@ class RepresentationModelRuntimeHints implements RuntimeHintsRegistrar { public void registerHints(RuntimeHints hints, ClassLoader classLoader) { var reflection = hints.reflection(); - var entityModelAndNested = Arrays.stream(EntityModel.class.getNestMembers()); - Stream.concat(REPRESENTATION_MODELS.stream(), entityModelAndNested).forEach(it -> { // - reflection.registerType(it, // - MemberCategory.INVOKE_DECLARED_CONSTRUCTORS, MemberCategory.INVOKE_DECLARED_METHODS); - }); + AotUtils.getScanner(RepresentationModel.class.getPackageName()) // + .findClasses() // + .forEach(it -> reflection.registerType(it, // + MemberCategory.INVOKE_DECLARED_CONSTRUCTORS, // + MemberCategory.INVOKE_DECLARED_METHODS)); } } diff --git a/src/main/java/org/springframework/hateoas/aot/HypermediaTypeAotProcessor.java b/src/main/java/org/springframework/hateoas/aot/HypermediaTypeAotProcessor.java index a18ef599..e572dfe3 100644 --- a/src/main/java/org/springframework/hateoas/aot/HypermediaTypeAotProcessor.java +++ b/src/main/java/org/springframework/hateoas/aot/HypermediaTypeAotProcessor.java @@ -17,6 +17,7 @@ package org.springframework.hateoas.aot; import java.io.IOException; import java.util.Arrays; +import java.util.Comparator; import java.util.HashSet; import java.util.List; import java.util.Set; @@ -28,22 +29,19 @@ import org.slf4j.LoggerFactory; import org.springframework.aot.generate.GenerationContext; import org.springframework.aot.hint.MemberCategory; import org.springframework.aot.hint.TypeReference; -import org.springframework.beans.factory.annotation.AnnotatedBeanDefinition; import org.springframework.beans.factory.aot.BeanRegistrationAotContribution; import org.springframework.beans.factory.aot.BeanRegistrationAotProcessor; import org.springframework.beans.factory.aot.BeanRegistrationCode; -import org.springframework.beans.factory.config.BeanDefinition; import org.springframework.beans.factory.support.RegisteredBean; -import org.springframework.context.annotation.ClassPathScanningCandidateComponentProvider; import org.springframework.core.annotation.AnnotatedElementUtils; import org.springframework.core.annotation.MergedAnnotation; import org.springframework.core.type.classreading.MetadataReader; import org.springframework.core.type.classreading.MetadataReaderFactory; import org.springframework.core.type.filter.TypeFilter; +import org.springframework.hateoas.aot.AotUtils.FullTypeScanner; import org.springframework.hateoas.config.EnableHypermediaSupport; import org.springframework.hateoas.config.EnableHypermediaSupport.HypermediaType; import org.springframework.util.Assert; -import org.springframework.util.ClassUtils; /** * A {@link BeanRegistrationAotProcessor} to register types that will be rendered by Jackson for reflection. The @@ -75,7 +73,7 @@ class HypermediaTypeAotProcessor implements BeanRegistrationAotProcessor { var mediaTypePackages = Stream.concat(fromConfig, Stream.of("alps", "problem")) .map("org.springframework.hateoas.mediatype."::concat); - var packagesToScan = Stream.concat(Stream.of("org.springframework.hateoas"), mediaTypePackages).toList(); + var packagesToScan = mediaTypePackages.toList(); return packagesToScan.isEmpty() ? null : new MediaTypeReflectionAotContribution(packagesToScan); } @@ -118,68 +116,21 @@ class HypermediaTypeAotProcessor implements BeanRegistrationAotProcessor { packagesSeen.add(it); // Register RepresentationModel types for full reflection - FullTypeScanner provider = new FullTypeScanner(); - provider.addIncludeFilter(new JacksonAnnotationPresentFilter()); - provider.addIncludeFilter(new JacksonSuperTypeFilter()); - - // Add filter to limit scan to sole package, not nested ones - provider.addExcludeFilter(new EnforcedPackageFilter(it)); + FullTypeScanner provider = AotUtils.getScanner(it, // + new JacksonAnnotationPresentFilter(), // + new JacksonSuperTypeFilter()); LOGGER.info("Registering Spring HATEOAS types in {} for reflection.", it); - provider.findCandidateComponents(it).stream() - .map(BeanDefinition::getBeanClassName) - .sorted() - .peek(type -> LOGGER.debug("> {}", type)) - .map(TypeReference::of) + provider.findClasses() + .sorted(Comparator.comparing(TypeReference::getName)) + .peek(type -> LOGGER.debug("> {}", type.getName())) .forEach(reference -> reflection.registerType(reference, // MemberCategory.INVOKE_DECLARED_CONSTRUCTORS, MemberCategory.INVOKE_DECLARED_METHODS)); }); } } - static class FullTypeScanner extends ClassPathScanningCandidateComponentProvider { - - public FullTypeScanner() { - super(false); - } - - /* - * (non-Javadoc) - * @see org.springframework.context.annotation.ClassPathScanningCandidateComponentProvider#isCandidateComponent(org.springframework.beans.factory.annotation.AnnotatedBeanDefinition) - */ - @Override - protected boolean isCandidateComponent(AnnotatedBeanDefinition beanDefinition) { - return true; - } - } - - /** - * A {@link TypeFilter} to only match types outside the configured package. Usually used as exclude filter - * to limit scans to not find nested packages. - * - * @author Oliver Drotbohm - */ - static class EnforcedPackageFilter implements TypeFilter { - - private final String referencePackage; - - public EnforcedPackageFilter(String referencePackage) { - this.referencePackage = referencePackage; - } - - /* - * (non-Javadoc) - * @see org.springframework.core.type.filter.TypeFilter#match(org.springframework.core.type.classreading.MetadataReader, org.springframework.core.type.classreading.MetadataReaderFactory) - */ - @Override - public boolean match(MetadataReader metadataReader, MetadataReaderFactory metadataReaderFactory) - throws IOException { - return !referencePackage - .equals(ClassUtils.getPackageName(metadataReader.getClassMetadata().getClassName())); - } - } - static abstract class TraversingTypeFilter implements TypeFilter { /* diff --git a/src/main/resources/META-INF/spring/aot.factories b/src/main/resources/META-INF/spring/aot.factories index 821b0b39..42b6d5ab 100644 --- a/src/main/resources/META-INF/spring/aot.factories +++ b/src/main/resources/META-INF/spring/aot.factories @@ -4,4 +4,4 @@ org.springframework.beans.factory.aot.BeanRegistrationAotProcessor=\ org.springframework.hateoas.aot.RepresentationModelAssemblerAotProcessor org.springframework.aot.hint.RuntimeHintsRegistrar=\ - org.springframework.hateoas.aot.RepresentationModelRuntimeHints + org.springframework.hateoas.aot.HateoasTypesRuntimeHints diff --git a/src/test/java/org/springframework/hateoas/aot/AotUtilsUnitTests.java b/src/test/java/org/springframework/hateoas/aot/AotUtilsUnitTests.java new file mode 100644 index 00000000..4af434b0 --- /dev/null +++ b/src/test/java/org/springframework/hateoas/aot/AotUtilsUnitTests.java @@ -0,0 +1,43 @@ +/* + * Copyright 2023 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 + * + * https://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.hateoas.aot; + +import static org.assertj.core.api.Assertions.*; + +import org.junit.jupiter.api.Test; +import org.springframework.aot.hint.TypeReference; +import org.springframework.hateoas.Link; +import org.springframework.hateoas.RepresentationModel; + +/** + * Unit tests for {@link AotUtils}. + * + * @author Oliver Drotbohm + */ +class AotUtilsUnitTests { + + @Test // GH-1981 + void findsTypesInPackage() { + + var scanner = AotUtils.getScanner(Link.class.getPackageName()); + + assertThat(scanner.findClasses()) + .extracting(TypeReference::getName) + .contains(Link.class.getName(), // + RepresentationModel.class.getName(), + "org.springframework.hateoas.EntityModel$MapSuppressingUnwrappingSerializer"); + } +} diff --git a/src/test/java/org/springframework/hateoas/aot/RepresentationModelRuntimeHintsUnitTests.java b/src/test/java/org/springframework/hateoas/aot/HateoasTypesRuntimeHintsUnitTests.java similarity index 78% rename from src/test/java/org/springframework/hateoas/aot/RepresentationModelRuntimeHintsUnitTests.java rename to src/test/java/org/springframework/hateoas/aot/HateoasTypesRuntimeHintsUnitTests.java index 94c0715b..463e29e1 100644 --- a/src/test/java/org/springframework/hateoas/aot/RepresentationModelRuntimeHintsUnitTests.java +++ b/src/test/java/org/springframework/hateoas/aot/HateoasTypesRuntimeHintsUnitTests.java @@ -21,18 +21,20 @@ import org.junit.jupiter.api.Test; import org.springframework.aot.hint.RuntimeHints; import org.springframework.aot.hint.TypeHint; import org.springframework.aot.hint.TypeReference; +import org.springframework.hateoas.Link; +import org.springframework.hateoas.Links; /** * Unit tests for {@link RepresentationModelRuntimeHints}. * * @author Oliver Drotbohm */ -class RepresentationModelRuntimeHintsUnitTests { +class HateoasTypesRuntimeHintsUnitTests { @Test // GH-1981 - void registersHintsForMapSuppressingUnwrappingSerializer() { + void registersHintsForHateoasTypes() { - var registrar = new RepresentationModelRuntimeHints(); + var registrar = new HateoasTypesRuntimeHints(); var hints = new RuntimeHints(); registrar.registerHints(hints, getClass().getClassLoader()); @@ -40,6 +42,8 @@ class RepresentationModelRuntimeHintsUnitTests { assertThat(hints.reflection().typeHints()) .extracting(TypeHint::getType) .extracting(TypeReference::getSimpleName) - .contains("MapSuppressingUnwrappingSerializer"); + .contains("MapSuppressingUnwrappingSerializer", // + Link.class.getSimpleName(), // + Links.class.getSimpleName()); } }