Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
/*
* Copyright 2012-present 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.boot.testcontainers.service.connection;

import org.jspecify.annotations.Nullable;

import org.springframework.beans.BeansException;
import org.springframework.beans.factory.BeanFactory;
import org.springframework.beans.factory.BeanFactoryAware;
import org.springframework.beans.factory.config.ConfigurableListableBeanFactory;
import org.springframework.boot.autoconfigure.service.connection.ConnectionDetails;
import org.springframework.boot.autoconfigure.service.connection.ConnectionDetailsFactories;
import org.springframework.util.Assert;

/**
* Factory used to create connection details from a container bean at runtime.
*
* @author Goutam Adwant
*/
class ConnectionDetailsBeanFactory implements BeanFactoryAware {

private @Nullable ConfigurableListableBeanFactory beanFactory;

private @Nullable ConnectionDetailsFactories connectionDetailsFactories;

@Override
public void setBeanFactory(BeanFactory beanFactory) throws BeansException {
Assert.isInstanceOf(ConfigurableListableBeanFactory.class, beanFactory);
ConfigurableListableBeanFactory listableBeanFactory = (ConfigurableListableBeanFactory) beanFactory;
this.beanFactory = listableBeanFactory;
this.connectionDetailsFactories = new ConnectionDetailsFactories(listableBeanFactory.getBeanClassLoader());
}

ConnectionDetails getConnectionDetails(String beanName, Class<?> connectionDetailsType) {
ConfigurableListableBeanFactory beanFactory = this.beanFactory;
Assert.state(beanFactory != null, "BeanFactory has not been set");
ConnectionDetailsFactories connectionDetailsFactories = this.connectionDetailsFactories;
Assert.state(connectionDetailsFactories != null, "ConnectionDetailsFactories has not been set");
for (ContainerConnectionSource<?> source : ServiceConnectionAutoConfigurationRegistrar.getSources(beanFactory,
beanName)) {
ConnectionDetails connectionDetails = connectionDetailsFactories.getConnectionDetails(source, true)
.get(connectionDetailsType);
if (connectionDetails != null) {
return connectionDetails;
}
}
throw new IllegalStateException("No connection details of type '%s' found for container bean '%s'"
.formatted(connectionDetailsType.getName(), beanName));
}

}
Original file line number Diff line number Diff line change
Expand Up @@ -49,11 +49,14 @@
* @author Moritz Halbritter
* @author Andy Wilkinson
* @author Phillip Webb
* @author Goutam Adwant
*/
class ConnectionDetailsRegistrar {

private static final Log logger = LogFactory.getLog(ConnectionDetailsRegistrar.class);

private static final String CONNECTION_DETAILS_BEAN_FACTORY = ConnectionDetailsBeanFactory.class.getName();

private final ListableBeanFactory beanFactory;

private final ConnectionDetailsFactories connectionDetailsFactories;
Expand Down Expand Up @@ -104,14 +107,33 @@ private <T> void registerBeanDefinition(BeanDefinitionRegistry registry, Contain
ContainerImageMetadata containerMetadata = new ContainerImageMetadata(source.getContainerImageName());
String beanName = getBeanName(source, connectionDetails);
Class<T> beanType = (Class<T>) connectionDetails.getClass();
Supplier<T> beanSupplier = () -> (T) connectionDetails;
logger.debug(LogMessage.of(() -> "Registering '%s' for %s".formatted(beanName, source)));
RootBeanDefinition beanDefinition = new RootBeanDefinition(beanType, beanSupplier);
beanDefinition.setAttribute(ServiceConnection.class.getName(), true);
RootBeanDefinition beanDefinition;
if (source.getOrigin() instanceof BeanOrigin) {
registerConnectionDetailsBeanFactory(registry);
beanDefinition = new RootBeanDefinition();
beanDefinition.setTargetType(beanType);
beanDefinition.setFactoryBeanName(CONNECTION_DETAILS_BEAN_FACTORY);
beanDefinition.setFactoryMethodName("getConnectionDetails");
beanDefinition.getConstructorArgumentValues().addIndexedArgumentValue(0, source.getBeanNameSuffix());
beanDefinition.getConstructorArgumentValues().addIndexedArgumentValue(1, connectionDetailsType);
}
else {
Supplier<T> beanSupplier = () -> (T) connectionDetails;
beanDefinition = new RootBeanDefinition(beanType, beanSupplier);
beanDefinition.setAttribute(ServiceConnection.class.getName(), true);
}
containerMetadata.addTo(beanDefinition);
registry.registerBeanDefinition(beanName, beanDefinition);
}

private void registerConnectionDetailsBeanFactory(BeanDefinitionRegistry registry) {
if (!registry.containsBeanDefinition(CONNECTION_DETAILS_BEAN_FACTORY)) {
registry.registerBeanDefinition(CONNECTION_DETAILS_BEAN_FACTORY,
new RootBeanDefinition(ConnectionDetailsBeanFactory.class));
}
}

private String getBeanName(ContainerConnectionSource<?> source, ConnectionDetails connectionDetails) {
List<String> parts = new ArrayList<>();
parts.add(ClassUtils.getShortNameAsProperty(connectionDetails.getClass()));
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,9 @@

package org.springframework.boot.testcontainers.service.connection;

import java.util.ArrayList;
import java.util.LinkedHashSet;
import java.util.List;
import java.util.Set;

import org.jspecify.annotations.Nullable;
Expand All @@ -28,6 +30,7 @@
import org.springframework.beans.factory.config.BeanDefinition;
import org.springframework.beans.factory.config.ConfigurableListableBeanFactory;
import org.springframework.beans.factory.support.BeanDefinitionRegistry;
import org.springframework.beans.factory.support.RootBeanDefinition;
import org.springframework.boot.autoconfigure.service.connection.ConnectionDetailsFactories;
import org.springframework.boot.origin.Origin;
import org.springframework.boot.testcontainers.beans.TestcontainerBeanDefinition;
Expand All @@ -44,6 +47,7 @@
*
* @author Phillip Webb
* @author Daeho Kwon
* @author Goutam Adwant
*/
class ServiceConnectionAutoConfigurationRegistrar implements ImportBeanDefinitionRegistrar {

Expand All @@ -64,18 +68,24 @@ private void registerBeanDefinitions(ConfigurableListableBeanFactory beanFactory
ConnectionDetailsRegistrar registrar = new ConnectionDetailsRegistrar(beanFactory,
new ConnectionDetailsFactories(null));
for (String beanName : beanFactory.getBeanNamesForType(Container.class)) {
BeanDefinition beanDefinition = getBeanDefinition(beanFactory, beanName);
MergedAnnotations annotations = getAnnotations(beanDefinition);
for (ServiceConnection serviceConnection : getServiceConnections(beanFactory, beanName, annotations)) {
ContainerConnectionSource<?> source = createSource(beanFactory, beanName, beanDefinition, annotations,
serviceConnection);
for (ContainerConnectionSource<?> source : getSources(beanFactory, beanName)) {
registrar.registerBeanDefinitions(registry, source);
}
}
}

private Set<ServiceConnection> getServiceConnections(ConfigurableListableBeanFactory beanFactory, String beanName,
@Nullable MergedAnnotations annotations) {
static List<ContainerConnectionSource<?>> getSources(ConfigurableListableBeanFactory beanFactory, String beanName) {
BeanDefinition beanDefinition = getBeanDefinition(beanFactory, beanName);
MergedAnnotations annotations = getAnnotations(beanDefinition);
List<ContainerConnectionSource<?>> sources = new ArrayList<>();
for (ServiceConnection serviceConnection : getServiceConnections(beanFactory, beanName, annotations)) {
sources.add(createSource(beanFactory, beanName, beanDefinition, annotations, serviceConnection));
}
return List.copyOf(sources);
}

private static Set<ServiceConnection> getServiceConnections(ConfigurableListableBeanFactory beanFactory,
String beanName, @Nullable MergedAnnotations annotations) {
Set<ServiceConnection> serviceConnections = beanFactory.findAllAnnotationsOnBean(beanName,
ServiceConnection.class, false);
if (annotations != null) {
Expand All @@ -87,7 +97,8 @@ private Set<ServiceConnection> getServiceConnections(ConfigurableListableBeanFac
return serviceConnections;
}

private @Nullable BeanDefinition getBeanDefinition(ConfigurableListableBeanFactory beanFactory, String beanName) {
private static @Nullable BeanDefinition getBeanDefinition(ConfigurableListableBeanFactory beanFactory,
String beanName) {
try {
return beanFactory.getBeanDefinition(beanName);
}
Expand All @@ -96,10 +107,14 @@ private Set<ServiceConnection> getServiceConnections(ConfigurableListableBeanFac
}
}

private @Nullable MergedAnnotations getAnnotations(@Nullable BeanDefinition beanDefinition) {
private static @Nullable MergedAnnotations getAnnotations(@Nullable BeanDefinition beanDefinition) {
if (beanDefinition instanceof TestcontainerBeanDefinition testcontainerBeanDefinition) {
return testcontainerBeanDefinition.getAnnotations();
}
if (beanDefinition instanceof RootBeanDefinition rootBeanDefinition
&& rootBeanDefinition.getResolvedFactoryMethod() != null) {
return MergedAnnotations.from(rootBeanDefinition.getResolvedFactoryMethod());
}
if (beanDefinition instanceof AnnotatedBeanDefinition annotatedBeanDefinition) {
MethodMetadata metadata = annotatedBeanDefinition.getFactoryMethodMetadata();
return (metadata != null) ? metadata.getAnnotations() : null;
Expand All @@ -108,7 +123,7 @@ private Set<ServiceConnection> getServiceConnections(ConfigurableListableBeanFac
}

@SuppressWarnings("unchecked")
private <C extends Container<?>> ContainerConnectionSource<C> createSource(
private static <C extends Container<?>> ContainerConnectionSource<C> createSource(
ConfigurableListableBeanFactory beanFactory, String beanName, @Nullable BeanDefinition beanDefinition,
@Nullable MergedAnnotations annotations, ServiceConnection serviceConnection) {
Origin origin = new BeanOrigin(beanName, beanDefinition);
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,121 @@
/*
* Copyright 2012-present 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.boot.testcontainers.service.connection;

import java.util.stream.Stream;

import org.junit.jupiter.api.Test;
import org.testcontainers.postgresql.PostgreSQLContainer;

import org.springframework.aot.AotDetector;
import org.springframework.aot.generate.InMemoryGeneratedFiles;
import org.springframework.aot.test.generate.CompilerFiles;
import org.springframework.boot.autoconfigure.ImportAutoConfiguration;
import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.boot.test.context.SpringBootTest.WebEnvironment;
import org.springframework.context.ApplicationContextInitializer;
import org.springframework.context.ConfigurableApplicationContext;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.core.test.tools.CompileWithForkedClassLoader;
import org.springframework.core.test.tools.TestCompiler;
import org.springframework.test.context.BootstrapUtils;
import org.springframework.test.context.MergedContextConfiguration;
import org.springframework.test.context.TestContextBootstrapper;
import org.springframework.test.context.aot.AotContextLoader;
import org.springframework.test.context.aot.AotTestContextInitializers;
import org.springframework.test.context.aot.TestContextAotGenerator;
import org.springframework.test.util.ReflectionTestUtils;
import org.springframework.util.ClassUtils;
import org.springframework.util.function.ThrowingConsumer;

import static org.assertj.core.api.Assertions.assertThat;
import static org.mockito.Mockito.mock;

/**
* Tests for {@link ServiceConnection} when used in AOT mode.
*
* @author Goutam Adwant
*/
@CompileWithForkedClassLoader
class ServiceConnectionAotTests {

@Test
void serviceConnectionOnBeanMethodIsAvailableAtAotRuntime() {
InMemoryGeneratedFiles generatedFiles = new InMemoryGeneratedFiles();
TestContextAotGenerator generator = new TestContextAotGenerator(generatedFiles);
Class<?> testClass = ExampleTest.class;
generator.processAheadOfTime(Stream.of(testClass));
TestCompiler.forSystem()
.withCompilerOptions("-Xlint:deprecation,removal", "-Werror")
.with(CompilerFiles.from(generatedFiles))
.compile(ThrowingConsumer.of((compiled) -> assertCompiledTest(testClass)));
}

private void assertCompiledTest(Class<?> testClass) throws Exception {
try {
System.setProperty(AotDetector.AOT_ENABLED, "true");
resetAotClasses();
AotTestContextInitializers aotContextInitializers = new AotTestContextInitializers();
TestContextBootstrapper testContextBootstrapper = BootstrapUtils.resolveTestContextBootstrapper(testClass);
MergedContextConfiguration mergedConfig = testContextBootstrapper.buildMergedContextConfiguration();
ApplicationContextInitializer<ConfigurableApplicationContext> contextInitializer = aotContextInitializers
.getContextInitializer(testClass);
assertThat(contextInitializer).isNotNull();
try (ConfigurableApplicationContext context = (ConfigurableApplicationContext) ((AotContextLoader) mergedConfig
.getContextLoader()).loadContextForAotRuntime(mergedConfig, contextInitializer)) {
assertThat(context.getBeansOfType(DatabaseConnectionDetails.class)).hasSize(1);
ContainerConnectionDetailsFactory.ContainerConnectionDetails<?> connectionDetails = (ContainerConnectionDetailsFactory.ContainerConnectionDetails<?>) context
.getBean(DatabaseConnectionDetails.class);
assertThat(connectionDetails.hasAnnotation(Ssl.class)).isTrue();
}
}
finally {
System.clearProperty(AotDetector.AOT_ENABLED);
resetAotClasses();
}
}

private void resetAotClasses() {
reset("org.springframework.test.context.aot.AotTestAttributesFactory");
reset("org.springframework.test.context.aot.AotTestContextInitializersFactory");
}

private void reset(String className) {
Class<?> targetClass = ClassUtils.resolveClassName(className, null);
ReflectionTestUtils.invokeMethod(targetClass, "reset");
}

@SpringBootTest(classes = ContainerConfiguration.class, webEnvironment = WebEnvironment.NONE)
static class ExampleTest {

}

@Configuration(proxyBeanMethods = false)
@ImportAutoConfiguration(ServiceConnectionAutoConfiguration.class)
static class ContainerConfiguration {

@Bean
@ServiceConnection
@Ssl
PostgreSQLContainer postgresContainer() {
return mock();
}

}

}