statements;
try {
- var sql = getFile(path);
- var split = sql.split(";");
- for (var exec : split) {
- if (exec != null && !exec.trim().isEmpty()) {
- try (var statement = connection.createStatement()) {
- statement.execute(exec);
- }
- }
+ statements = SqlScriptParser.splitStatements(getFile(path));
+ } catch (IllegalArgumentException ex) {
+ throw new ClickhouseStartingException("Invalid SQL migration " + path, ex);
+ }
+ for (int index = 0; index < statements.size(); index++) {
+ try (var statement = connection.createStatement()) {
+ statement.execute(statements.get(index));
+ } catch (SQLException ex) {
+ throw new ClickhouseStartingException(
+ "Error while executing statement " + (index + 1) + " from migration " + path,
+ ex);
}
- } catch (SQLException e) {
- log.error("Error when execAllInFile path: {}", path);
- throw new ClickhouseStartingException(String.format("Error when execAllInFile path: %s", path), e);
}
}
@@ -96,9 +99,17 @@ private String getFile(String fileName) {
"Migration file not found: " + fileName,
new FileNotFoundException(fileName)))) {
return IOUtils.toString(inputStream, StandardCharsets.UTF_8);
- } catch (IOException e) {
- log.error("Error when getFile e: ", e);
- throw new ClickhouseStartingException("Error when reading migration file: " + fileName, e);
+ } catch (IOException ex) {
+ throw new ClickhouseStartingException("Error while reading migration file " + fileName, ex);
+ }
+ }
+
+ private static String validateIdentifier(String identifier) {
+ if (identifier == null || !SAFE_IDENTIFIER.matcher(identifier).matches()) {
+ throw new IllegalArgumentException(
+ "Unsafe ClickHouse database identifier: " + identifier +
+ ". Only letters, digits and underscores are supported");
}
+ return identifier;
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/clickhouse/ClickhouseTestcontainerExtension.java b/src/main/java/dev/vality/testcontainers/annotations/clickhouse/ClickhouseTestcontainerExtension.java
index c1fe0dcc..c7b974e8 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/clickhouse/ClickhouseTestcontainerExtension.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/clickhouse/ClickhouseTestcontainerExtension.java
@@ -1,7 +1,8 @@
package dev.vality.testcontainers.annotations.clickhouse;
import dev.vality.testcontainers.annotations.util.GenericContainerUtil;
-import lombok.extern.slf4j.Slf4j;
+import dev.vality.testcontainers.annotations.util.SharedTestResourceLock;
+import dev.vality.testcontainers.annotations.util.TestExecutionLock;
import org.junit.jupiter.api.extension.AfterAllCallback;
import org.junit.jupiter.api.extension.BeforeAllCallback;
import org.junit.jupiter.api.extension.BeforeEachCallback;
@@ -12,114 +13,178 @@
import org.springframework.test.context.ContextConfigurationAttributes;
import org.springframework.test.context.ContextCustomizer;
import org.springframework.test.context.ContextCustomizerFactory;
+import org.springframework.test.context.MergedContextConfiguration;
+import java.util.Arrays;
import java.util.List;
import java.util.Optional;
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.ConcurrentMap;
-/**
- * {@code @ClickhouseTestcontainerExtension} инициализирует тестконтейнер из {@link ClickhouseTestcontainerFactory},
- * настраивает, стартует, валидирует и останавливает
- *
{@link ClickhouseTestcontainerExtension.ClickhouseTestcontainerContextCustomizerFactory}
- * Инициализация настроек контейнеров в спринговый контекст тестового приложения реализован
- * под капотом аннотаций, на уровне реализации интерфейса —
- * информация о настройках используемого тестконтейнера и передаваемые через параметры аннотации настройки
- * инициализируются через {@link TestPropertyValues} и сливаются с текущим получаемым контекстом
- * приложения {@link ConfigurableApplicationContext}
- *
Инициализация кастомизированных фабрик с инициализацией настроек осуществляется через описание бинов
- * в файле META-INF/spring.factories
- *
- * @see ClickhouseTestcontainerFactory ClickhouseTestcontainerFactory
- * @see ClickhouseTestcontainerExtension.ClickhouseTestcontainerContextCustomizerFactory ClickhouseTestcontainerContextCustomizerFactory
- * @see TestPropertyValues TestPropertyValues
- * @see ConfigurableApplicationContext ConfigurableApplicationContext
- * @see BeforeAllCallback BeforeAllCallback
- * @see AfterAllCallback AfterAllCallback
- */
-@Slf4j
public class ClickhouseTestcontainerExtension implements BeforeAllCallback, AfterAllCallback, BeforeEachCallback {
- private static final ThreadLocal THREAD_CONTAINER = new ThreadLocal<>();
+ private static final ConcurrentMap, ContainerReference> CONTAINERS = new ConcurrentHashMap<>();
@Override
public void beforeAll(ExtensionContext context) {
- if (findPrototypeAnnotation(context).isPresent()) {
- var annotation = findPrototypeAnnotation(context).get();
- var container = ClickhouseTestcontainerFactory.container(annotation.dbNameShouldBeDropped(),
- annotation.migrations());
- GenericContainerUtil.startContainer(container);
- THREAD_CONTAINER.set(container);
- } else if (findSingletonAnnotation(context).isPresent()) {
- var annotation = findSingletonAnnotation(context).get();
- var container = ClickhouseTestcontainerFactory.singletonContainer(annotation.dbNameShouldBeDropped(),
- annotation.migrations());
- if (!container.isRunning()) {
- GenericContainerUtil.startContainer(container);
- }
- THREAD_CONTAINER.set(container);
- }
+ getOrStart(context.getRequiredTestClass());
}
@Override
public void beforeEach(ExtensionContext context) {
- var container = THREAD_CONTAINER.get();
- if (container != null && container.isRunning()) {
- container.dropDatabase();
- container.appliedMigrations();
+ TestExecutionLock.acquire(context);
+ try {
+ var reference = getOrStart(context.getRequiredTestClass());
+ reference.container().dropDatabase();
+ reference.container().applyMigrations();
+ } catch (RuntimeException ex) {
+ TestExecutionLock.release(context);
+ throw ex;
}
}
@Override
public void afterAll(ExtensionContext context) {
- if (findPrototypeAnnotation(context).isPresent()) {
- var container = THREAD_CONTAINER.get();
- if (container != null && container.isRunning()) {
- container.stop();
+ var testClass = context.getRequiredTestClass();
+ var reference = CONTAINERS.remove(testClass);
+ try {
+ if (reference != null && !reference.singleton()) {
+ reference.container().stop();
+ }
+ } finally {
+ if (reference != null && reference.singleton()) {
+ SharedTestResourceLock.release(testClass);
}
- THREAD_CONTAINER.remove();
- } else if (findSingletonAnnotation(context).isPresent()) {
- THREAD_CONTAINER.remove();
}
}
- private static Optional findPrototypeAnnotation(ExtensionContext context) {
- return AnnotationSupport.findAnnotation(context.getTestClass(), ClickhouseTestcontainer.class);
+ private static ContainerReference getOrStart(Class> testClass) {
+ return CONTAINERS.computeIfAbsent(testClass, ClickhouseTestcontainerExtension::createAndStart);
}
- private static Optional findPrototypeAnnotation(Class> testClass) {
- return AnnotationSupport.findAnnotation(testClass, ClickhouseTestcontainer.class);
+ private static ContainerReference createAndStart(Class> testClass) {
+ var prototype = findPrototypeAnnotation(testClass);
+ var singletonAnnotation = findSingletonAnnotation(testClass);
+ if (prototype.isEmpty() && singletonAnnotation.isEmpty()) {
+ throw new IllegalStateException("ClickHouse test annotation not found");
+ }
+
+ var singleton = singletonAnnotation.isPresent();
+ if (singleton) {
+ SharedTestResourceLock.acquire(testClass);
+ }
+ try {
+ var container = prototype
+ .map(annotation -> ClickhouseTestcontainerFactory.container(
+ annotation.dbNameShouldBeDropped(),
+ annotation.migrations()))
+ .orElseGet(() -> {
+ var annotation = singletonAnnotation.orElseThrow();
+ return ClickhouseTestcontainerFactory.singletonContainer(
+ annotation.dbNameShouldBeDropped(),
+ annotation.migrations());
+ });
+ try {
+ GenericContainerUtil.startContainer(container);
+ container.dropDatabase();
+ container.applyMigrations();
+ return new ContainerReference(container, singleton);
+ } catch (RuntimeException ex) {
+ try {
+ if (singleton) {
+ ClickhouseTestcontainerFactory.discardSingleton(container);
+ } else {
+ container.stop();
+ }
+ } catch (RuntimeException stopException) {
+ ex.addSuppressed(stopException);
+ }
+ throw ex;
+ }
+ } catch (RuntimeException ex) {
+ if (singleton) {
+ SharedTestResourceLock.release(testClass);
+ }
+ throw ex;
+ }
}
- private static Optional findSingletonAnnotation(ExtensionContext context) {
- return AnnotationSupport.findAnnotation(context.getTestClass(), ClickhouseTestcontainerSingleton.class);
+ private static Optional findPrototypeAnnotation(Class> testClass) {
+ return AnnotationSupport.findAnnotation(testClass, ClickhouseTestcontainer.class);
}
private static Optional findSingletonAnnotation(Class> testClass) {
return AnnotationSupport.findAnnotation(testClass, ClickhouseTestcontainerSingleton.class);
}
+ private static final class ContainerReference {
+
+ private final ClickhouseContainerExtension container;
+ private final boolean singleton;
+
+ private ContainerReference(ClickhouseContainerExtension container, boolean singleton) {
+ this.container = container;
+ this.singleton = singleton;
+ }
+
+ private ClickhouseContainerExtension container() {
+ return container;
+ }
+
+ private boolean singleton() {
+ return singleton;
+ }
+
+ }
+
public static class ClickhouseTestcontainerContextCustomizerFactory implements ContextCustomizerFactory {
@Override
public ContextCustomizer createContextCustomizer(
Class> testClass,
List configAttributes) {
- return (context, mergedConfig) -> {
- if (findPrototypeAnnotation(testClass).isPresent()) {
- init(context, findPrototypeAnnotation(testClass).get().properties());
- } else if (findSingletonAnnotation(testClass).isPresent()) {
- init(context, findSingletonAnnotation(testClass).get().properties());
- }
- };
+ var prototype = findPrototypeAnnotation(testClass);
+ if (prototype.isPresent()) {
+ var annotation = prototype.get();
+ return new ClickhouseContextCustomizer(
+ testClass,
+ false,
+ annotation.dbNameShouldBeDropped(),
+ List.copyOf(Arrays.asList(annotation.migrations())),
+ List.copyOf(Arrays.asList(annotation.properties())));
+ }
+ var singleton = findSingletonAnnotation(testClass);
+ if (singleton.isPresent()) {
+ var annotation = singleton.get();
+ return new ClickhouseContextCustomizer(
+ testClass,
+ true,
+ annotation.dbNameShouldBeDropped(),
+ List.copyOf(Arrays.asList(annotation.migrations())),
+ List.copyOf(Arrays.asList(annotation.properties())));
+ }
+ return null;
}
+ }
- private void init(ConfigurableApplicationContext context, String[] properties) {
- var container = THREAD_CONTAINER.get();
+ private record ClickhouseContextCustomizer(
+ Class> testClass,
+ boolean singleton,
+ String databaseName,
+ List migrations,
+ List properties) implements ContextCustomizer {
+
+ @Override
+ public void customizeContext(
+ ConfigurableApplicationContext context,
+ MergedContextConfiguration mergedConfig) {
+ var container = getOrStart(testClass).container();
TestPropertyValues.of(
"clickhouse.db.url=" + container.getJdbcUrl(),
"clickhouse.db.user=" + container.getUsername(),
"clickhouse.db.username=" + container.getUsername(),
"clickhouse.db.password=" + container.getPassword())
- .and(properties)
+ .and(properties.toArray(String[]::new))
.applyTo(context);
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/clickhouse/ClickhouseTestcontainerFactory.java b/src/main/java/dev/vality/testcontainers/annotations/clickhouse/ClickhouseTestcontainerFactory.java
index 4a30dea4..1920afa4 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/clickhouse/ClickhouseTestcontainerFactory.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/clickhouse/ClickhouseTestcontainerFactory.java
@@ -1,51 +1,65 @@
package dev.vality.testcontainers.annotations.clickhouse;
+import dev.vality.testcontainers.annotations.util.ContainerShutdownRegistry;
import lombok.AccessLevel;
import lombok.NoArgsConstructor;
-import lombok.Synchronized;
-import lombok.extern.slf4j.Slf4j;
+
+import java.util.Arrays;
+import java.util.List;
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.ConcurrentMap;
/**
- * Фабрика по созданию контейнеров
- * {@link #create(String, String[])} создает экземпляр тестконтейнера
- *
{@link #getOrCreateSingletonContainer(String, String[])} создает синглтон тестконтейнера
- *
- * @see ClickhouseTestcontainerExtension ClickhouseTestcontainerExtension
+ * Фабрика по созданию контейнеров.
*/
-@Slf4j
@NoArgsConstructor(access = AccessLevel.PRIVATE)
public class ClickhouseTestcontainerFactory {
- private ClickhouseContainerExtension clickHouseContainer;
+ private final ConcurrentMap singletonContainers = new ConcurrentHashMap<>();
public static ClickhouseContainerExtension container(String databaseName, String[] migrations) {
- return instance().create(databaseName, migrations);
+ return instance().create(Config.of(databaseName, migrations));
}
public static ClickhouseContainerExtension singletonContainer(String databaseName, String[] migrations) {
- return instance().getOrCreateSingletonContainer(databaseName, migrations);
+ var config = Config.of(databaseName, migrations);
+ return instance().singletonContainers.computeIfAbsent(
+ config,
+ key -> ContainerShutdownRegistry.register(instance().create(key)));
+ }
+
+ static void discardSingleton(ClickhouseContainerExtension container) {
+ if (instance().singletonContainers.values().remove(container)) {
+ ContainerShutdownRegistry.unregister(container);
+ container.stop();
+ }
}
private static ClickhouseTestcontainerFactory instance() {
return SingletonHolder.INSTANCE;
}
- @Synchronized
- private ClickhouseContainerExtension getOrCreateSingletonContainer(String databaseName, String[] migrations) {
- if (clickHouseContainer != null) {
- return clickHouseContainer;
- }
- clickHouseContainer = create(databaseName, migrations);
- return clickHouseContainer;
+ private ClickhouseContainerExtension create(Config config) {
+ return new ClickhouseContainerExtension(
+ config.databaseName(),
+ config.migrations().toArray(String[]::new));
}
- private ClickhouseContainerExtension create(String databaseName, String[] migrations) {
- return new ClickhouseContainerExtension(databaseName, migrations);
+ private record Config(String databaseName, List migrations) {
+
+ private static Config of(String databaseName, String[] migrations) {
+ if (databaseName == null || databaseName.isBlank()) {
+ throw new IllegalArgumentException("ClickHouse database name must not be blank");
+ }
+ var migrationList = migrations == null
+ ? List.of()
+ : Arrays.stream(migrations).map(String::trim).toList();
+ return new Config(databaseName.trim(), List.copyOf(migrationList));
+ }
}
private static class SingletonHolder {
private static final ClickhouseTestcontainerFactory INSTANCE = new ClickhouseTestcontainerFactory();
-
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/clickhouse/SqlScriptParser.java b/src/main/java/dev/vality/testcontainers/annotations/clickhouse/SqlScriptParser.java
new file mode 100644
index 00000000..5b8b9341
--- /dev/null
+++ b/src/main/java/dev/vality/testcontainers/annotations/clickhouse/SqlScriptParser.java
@@ -0,0 +1,201 @@
+package dev.vality.testcontainers.annotations.clickhouse;
+
+import java.util.ArrayList;
+import java.util.List;
+
+final class SqlScriptParser {
+
+ private SqlScriptParser() {
+ }
+
+ static List splitStatements(String script) {
+ var parser = new Parser();
+ var index = 0;
+ while (index < script.length()) {
+ index = parser.consume(script, index);
+ }
+ return parser.finish();
+ }
+
+ private static String findDollarQuoteDelimiter(String script, int start) {
+ var end = script.indexOf('$', start + 1);
+ if (end < 0) {
+ return null;
+ }
+ var tag = script.substring(start + 1, end);
+ if (!tag.chars().allMatch(character -> Character.isLetterOrDigit(character) || character == '_')) {
+ return null;
+ }
+ return script.substring(start, end + 1);
+ }
+
+ private static void addStatement(List statements, StringBuilder current) {
+ var statement = current.toString().trim();
+ current.setLength(0);
+ if (!statement.isEmpty()) {
+ statements.add(statement);
+ }
+ }
+
+ private static final class Parser {
+
+ private final List statements = new ArrayList<>();
+ private final StringBuilder current = new StringBuilder();
+ private boolean singleQuoted;
+ private boolean doubleQuoted;
+ private boolean backtickQuoted;
+ private boolean lineComment;
+ private boolean blockComment;
+ private String dollarQuote;
+
+ private int consume(String script, int index) {
+ if (lineComment) {
+ return consumeLineComment(script, index);
+ }
+ if (blockComment) {
+ return consumeBlockComment(script, index);
+ }
+ if (dollarQuote != null) {
+ return consumeDollarQuote(script, index);
+ }
+ if (!isQuoted()) {
+ var nextIndex = consumeUnquotedToken(script, index);
+ if (nextIndex >= 0) {
+ return nextIndex;
+ }
+ }
+ return consumeText(script, index);
+ }
+
+ private int consumeLineComment(String script, int index) {
+ var currentChar = script.charAt(index);
+ if (currentChar == '\n') {
+ current.append(currentChar);
+ lineComment = false;
+ }
+ return index + 1;
+ }
+
+ private int consumeBlockComment(String script, int index) {
+ if (startsWith(script, index, "*/")) {
+ current.append(' ');
+ blockComment = false;
+ return index + 2;
+ }
+ return index + 1;
+ }
+
+ private int consumeDollarQuote(String script, int index) {
+ if (script.startsWith(dollarQuote, index)) {
+ current.append(dollarQuote);
+ var nextIndex = index + dollarQuote.length();
+ dollarQuote = null;
+ return nextIndex;
+ }
+ current.append(script.charAt(index));
+ return index + 1;
+ }
+
+ private int consumeUnquotedToken(String script, int index) {
+ if (startsWith(script, index, "--")) {
+ current.append(' ');
+ lineComment = true;
+ return index + 2;
+ }
+ if (startsWith(script, index, "/*")) {
+ current.append(' ');
+ blockComment = true;
+ return index + 2;
+ }
+ if (script.charAt(index) == '$') {
+ return consumeDollarQuoteStart(script, index);
+ }
+ if (script.charAt(index) == ';') {
+ addStatement(statements, current);
+ return index + 1;
+ }
+ return -1;
+ }
+
+ private int consumeDollarQuoteStart(String script, int index) {
+ var delimiter = findDollarQuoteDelimiter(script, index);
+ if (delimiter == null) {
+ return -1;
+ }
+ current.append(delimiter);
+ dollarQuote = delimiter;
+ return index + delimiter.length();
+ }
+
+ private int consumeText(String script, int index) {
+ var currentChar = script.charAt(index);
+ var nextChar = nextChar(script, index);
+ current.append(currentChar);
+ if (currentChar == '\\' && isQuoted() && nextChar != '\0') {
+ current.append(nextChar);
+ return index + 2;
+ }
+ return consumeQuote(currentChar, nextChar, index);
+ }
+
+ private int consumeQuote(char currentChar, char nextChar, int index) {
+ if (currentChar == '\'' && !doubleQuoted && !backtickQuoted) {
+ return consumeSingleQuote(nextChar, index);
+ }
+ if (currentChar == '"' && !singleQuoted && !backtickQuoted) {
+ return consumeDoubleQuote(nextChar, index);
+ }
+ if (currentChar == '`' && !singleQuoted && !doubleQuoted) {
+ return consumeBacktickQuote(nextChar, index);
+ }
+ return index + 1;
+ }
+
+ private int consumeSingleQuote(char nextChar, int index) {
+ if (singleQuoted && nextChar == '\'') {
+ current.append(nextChar);
+ return index + 2;
+ }
+ singleQuoted = !singleQuoted;
+ return index + 1;
+ }
+
+ private int consumeDoubleQuote(char nextChar, int index) {
+ if (doubleQuoted && nextChar == '"') {
+ current.append(nextChar);
+ return index + 2;
+ }
+ doubleQuoted = !doubleQuoted;
+ return index + 1;
+ }
+
+ private int consumeBacktickQuote(char nextChar, int index) {
+ if (backtickQuoted && nextChar == '`') {
+ current.append(nextChar);
+ return index + 2;
+ }
+ backtickQuoted = !backtickQuoted;
+ return index + 1;
+ }
+
+ private boolean isQuoted() {
+ return singleQuoted || doubleQuoted || backtickQuoted;
+ }
+
+ private List finish() {
+ if (isQuoted() || blockComment || dollarQuote != null) {
+ throw new IllegalArgumentException("SQL script contains an unterminated quoted value or comment");
+ }
+ addStatement(statements, current);
+ return List.copyOf(statements);
+ }
+
+ private static boolean startsWith(String script, int index, String value) {
+ return script.startsWith(value, index);
+ }
+
+ private static char nextChar(String script, int index) {
+ return index + 1 < script.length() ? script.charAt(index + 1) : '\0';
+ }
+ }
+}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/kafka/ApacheKafkaContainer.java b/src/main/java/dev/vality/testcontainers/annotations/kafka/ApacheKafkaContainer.java
index 4d8ae98e..3a0e19c5 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/kafka/ApacheKafkaContainer.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/kafka/ApacheKafkaContainer.java
@@ -1,12 +1,10 @@
package dev.vality.testcontainers.annotations.kafka;
import lombok.extern.slf4j.Slf4j;
-import org.testcontainers.containers.Network;
import org.testcontainers.kafka.KafkaContainer;
import org.testcontainers.utility.DockerImageName;
import java.util.List;
-import java.util.UUID;
import static dev.vality.testcontainers.annotations.util.SpringApplicationPropertiesLoader.loadDefaultLibraryProperty;
@@ -30,8 +28,6 @@ public ApacheKafkaContainer(List topics) {
withEnv("KAFKA_CFG_AUTO_CREATE_TOPICS_ENABLE", "false");
withEnv("KAFKA_AUTO_CREATE_TOPICS_ENABLE", "false");
withEnv("KAFKA_CFG_NODE_ID", "1");
- withNetworkAliases("kafka-" + UUID.randomUUID());
- withNetwork(Network.SHARED);
}
@Override
diff --git a/src/main/java/dev/vality/testcontainers/annotations/kafka/ConfluentKafkaContainer.java b/src/main/java/dev/vality/testcontainers/annotations/kafka/ConfluentKafkaContainer.java
index 7b7e2fee..d6d54297 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/kafka/ConfluentKafkaContainer.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/kafka/ConfluentKafkaContainer.java
@@ -1,11 +1,9 @@
package dev.vality.testcontainers.annotations.kafka;
import lombok.extern.slf4j.Slf4j;
-import org.testcontainers.containers.Network;
import org.testcontainers.utility.DockerImageName;
import java.util.List;
-import java.util.UUID;
import static dev.vality.testcontainers.annotations.util.SpringApplicationPropertiesLoader.loadDefaultLibraryProperty;
@@ -30,8 +28,6 @@ public ConfluentKafkaContainer(List topics) {
withEnv("KAFKA_CFG_AUTO_CREATE_TOPICS_ENABLE", "false");
withEnv("KAFKA_AUTO_CREATE_TOPICS_ENABLE", "false");
withEnv("KAFKA_CFG_NODE_ID", "1");
- withNetworkAliases("kafka-" + UUID.randomUUID());
- withNetwork(Network.SHARED);
}
@Override
diff --git a/src/main/java/dev/vality/testcontainers/annotations/kafka/EmbeddedKafkaTestContextCustomizerFactory.java b/src/main/java/dev/vality/testcontainers/annotations/kafka/EmbeddedKafkaTestContextCustomizerFactory.java
index 9cac00f3..dc692fbd 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/kafka/EmbeddedKafkaTestContextCustomizerFactory.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/kafka/EmbeddedKafkaTestContextCustomizerFactory.java
@@ -6,9 +6,10 @@
import org.springframework.test.context.ContextConfigurationAttributes;
import org.springframework.test.context.ContextCustomizer;
import org.springframework.test.context.ContextCustomizerFactory;
+import org.springframework.test.context.MergedContextConfiguration;
+import java.util.Arrays;
import java.util.List;
-import java.util.Optional;
public class EmbeddedKafkaTestContextCustomizerFactory implements ContextCustomizerFactory {
@@ -16,20 +17,24 @@ public class EmbeddedKafkaTestContextCustomizerFactory implements ContextCustomi
public ContextCustomizer createContextCustomizer(
Class> testClass,
List configAttributes) {
- return (context, mergedConfig) ->
- findAnnotation(testClass).ifPresent(annotation -> init(context, annotation));
+ return AnnotationSupport.findAnnotation(testClass, EmbeddedKafkaTest.class)
+ .map(annotation -> new EmbeddedKafkaContextCustomizer(
+ List.copyOf(Arrays.asList(annotation.properties()))))
+ .orElse(null);
}
- private Optional findAnnotation(Class> testClass) {
- return AnnotationSupport.findAnnotation(testClass, EmbeddedKafkaTest.class);
- }
+ private record EmbeddedKafkaContextCustomizer(List properties) implements ContextCustomizer {
- private void init(ConfigurableApplicationContext context, EmbeddedKafkaTest annotation) {
- TestPropertyValues.of(
- "spring.kafka.bootstrap-servers=${spring.embedded.kafka.brokers}",
- "kafka.bootstrap-servers=${spring.embedded.kafka.brokers}",
- "kafka.ssl.enabled=false")
- .and(annotation.properties())
- .applyTo(context);
+ @Override
+ public void customizeContext(
+ ConfigurableApplicationContext context,
+ MergedContextConfiguration mergedConfig) {
+ TestPropertyValues.of(
+ "spring.kafka.bootstrap-servers=${spring.embedded.kafka.brokers}",
+ "kafka.bootstrap-servers=${spring.embedded.kafka.brokers}",
+ "kafka.ssl.enabled=false")
+ .and(properties.toArray(String[]::new))
+ .applyTo(context);
+ }
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/kafka/EmbeddedKafkaTestExtension.java b/src/main/java/dev/vality/testcontainers/annotations/kafka/EmbeddedKafkaTestExtension.java
index 1b4fa670..405dbb79 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/kafka/EmbeddedKafkaTestExtension.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/kafka/EmbeddedKafkaTestExtension.java
@@ -1,6 +1,6 @@
package dev.vality.testcontainers.annotations.kafka;
-import lombok.SneakyThrows;
+import dev.vality.testcontainers.annotations.util.TestExecutionLock;
import org.apache.kafka.clients.admin.AdminClient;
import org.apache.kafka.clients.admin.AdminClientConfig;
import org.apache.kafka.clients.admin.RecordsToDelete;
@@ -15,7 +15,9 @@
import org.springframework.test.context.junit.jupiter.SpringExtension;
import java.util.*;
+import java.util.concurrent.ExecutionException;
import java.util.concurrent.TimeUnit;
+import java.util.concurrent.TimeoutException;
public class EmbeddedKafkaTestExtension implements BeforeEachCallback {
@@ -28,29 +30,42 @@ public void beforeEach(ExtensionContext context) {
.ifPresent(annotation -> cleanupTopics(context, annotation));
}
- @SneakyThrows
private void cleanupTopics(ExtensionContext context, EmbeddedKafkaTest annotation) {
+ var excludedTopics = Set.copyOf(Arrays.asList(annotation.excludeCleanupTopics()));
var topics = Arrays.stream(annotation.topics())
- .filter(topic -> !List.of(annotation.excludeCleanupTopics()).contains(topic))
+ .filter(topic -> !excludedTopics.contains(topic))
.toList();
if (topics.isEmpty()) {
return;
}
- var applicationContext = SpringExtension.getApplicationContext(context);
- var embeddedKafkaBroker = applicationContext.getBean(EmbeddedKafkaBroker.class);
- try (var adminClient = AdminClient.create(buildAdminProperties(embeddedKafkaBroker))) {
- var existingTopics = adminClient.listTopics().names().get(WAIT_TIMEOUT_SECONDS, TimeUnit.SECONDS);
- var missingTopics = topics.stream()
- .filter(topic -> !existingTopics.contains(topic))
- .toList();
- if (!missingTopics.isEmpty()) {
- embeddedKafkaBroker.addTopics(missingTopics.toArray(String[]::new));
- }
- var recordsToDelete = buildRecordsToDelete(embeddedKafkaBroker, topics);
- if (!recordsToDelete.isEmpty()) {
- adminClient.deleteRecords(recordsToDelete).all().get(WAIT_TIMEOUT_SECONDS, TimeUnit.SECONDS);
+ TestExecutionLock.acquire(context);
+ try {
+ var applicationContext = SpringExtension.getApplicationContext(context);
+ var embeddedKafkaBroker = applicationContext.getBean(EmbeddedKafkaBroker.class);
+ try (var adminClient = AdminClient.create(buildAdminProperties(embeddedKafkaBroker))) {
+ var existingTopics = adminClient.listTopics().names().get(WAIT_TIMEOUT_SECONDS, TimeUnit.SECONDS);
+ var missingTopics = topics.stream()
+ .filter(topic -> !existingTopics.contains(topic))
+ .toList();
+ if (!missingTopics.isEmpty()) {
+ embeddedKafkaBroker.addTopics(missingTopics.toArray(String[]::new));
+ }
+ var recordsToDelete = buildRecordsToDelete(embeddedKafkaBroker, topics);
+ if (!recordsToDelete.isEmpty()) {
+ adminClient.deleteRecords(recordsToDelete).all().get(WAIT_TIMEOUT_SECONDS, TimeUnit.SECONDS);
+ }
}
+ } catch (InterruptedException ex) {
+ Thread.currentThread().interrupt();
+ TestExecutionLock.release(context);
+ throw new IllegalStateException("Interrupted while cleaning embedded Kafka topics", ex);
+ } catch (ExecutionException | TimeoutException ex) {
+ TestExecutionLock.release(context);
+ throw new IllegalStateException("Unable to clean embedded Kafka topics", ex);
+ } catch (RuntimeException ex) {
+ TestExecutionLock.release(context);
+ throw ex;
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaContainerExtension.java b/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaContainerExtension.java
index 66307727..0d9607e5 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaContainerExtension.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaContainerExtension.java
@@ -12,14 +12,13 @@
import java.io.IOException;
import java.time.Duration;
+import java.util.Arrays;
import java.util.List;
import java.util.Properties;
+import java.util.Set;
import java.util.concurrent.ExecutionException;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.TimeoutException;
-import java.util.stream.Collectors;
-
-import static org.assertj.core.api.Assertions.assertThat;
public interface KafkaContainerExtension extends Startable, ContainerState {
@@ -34,76 +33,79 @@ public interface KafkaContainerExtension extends Startable, ContainerState {
String execInContainerKafkaTopicsListCommand();
default void createTopics(List excludedTopics) {
+ var topicsToEnsure = topics().stream()
+ .filter(topic -> !excludedTopics.contains(topic))
+ .toList();
+ if (topicsToEnsure.isEmpty()) {
+ return;
+ }
+
try (var admin = createAdminClient()) {
- var topics = topics().stream()
- .filter(topic -> !excludedTopics.contains(topic))
- .toList();
- if (topics.isEmpty()) {
- return;
- }
var existingTopics = admin.listTopics().names().get(WAIT_TIMEOUT, TimeUnit.SECONDS);
- var topicsToCreate = topics.stream()
+ var topicsToCreate = topicsToEnsure.stream()
.filter(topic -> !existingTopics.contains(topic))
.toList();
- if (topicsToCreate.isEmpty()) {
- return;
+ if (!topicsToCreate.isEmpty()) {
+ var newTopics = topicsToCreate.stream()
+ .map(topic -> new NewTopic(topic, 1, (short) 1))
+ .toList();
+ admin.createTopics(newTopics).all().get(WAIT_TIMEOUT, TimeUnit.SECONDS);
}
- var newTopics = topicsToCreate.stream()
- .map(topic -> new NewTopic(topic, 1, (short) 1))
- .peek(newTopic -> log.info(newTopic.toString()))
- .collect(Collectors.toList());
- var topicsResult = admin.createTopics(newTopics);
+
Awaitility.await()
.atMost(Duration.ofSeconds(WAIT_TIMEOUT))
- .pollInterval(Duration.ofSeconds(2))
- .untilAsserted(() -> topicsResult.all().get(1, TimeUnit.SECONDS));
- var topicsAfterCreate = admin.listTopics().names().get(WAIT_TIMEOUT, TimeUnit.SECONDS);
- log.info("Topics list from 'AdminClient' after [TOPICS CREATED]: {}", topicsAfterCreate);
- assertThat(topicsAfterCreate)
- .containsAll(topicsToCreate);
- var actual = execInContainerKafkaTopicsListCommand();
- assertThat(topicsToCreate.stream().allMatch(actual::contains))
- .isTrue();
+ .pollInterval(Duration.ofSeconds(1))
+ .until(() -> {
+ var actual = admin.listTopics().names().get(5, TimeUnit.SECONDS);
+ return actual.containsAll(topicsToEnsure);
+ });
+
+ var actual = admin.listTopics().names().get(WAIT_TIMEOUT, TimeUnit.SECONDS);
+ if (!actual.containsAll(topicsToEnsure)) {
+ throw new KafkaStartingException(
+ "Kafka topics were not created. Expected=" + topicsToEnsure + ", actual=" + actual,
+ null);
+ }
+ log.info("Kafka topics are ready: {}", topicsToEnsure);
} catch (ExecutionException | TimeoutException ex) {
- throw new KafkaStartingException("Error when topic creating, ", ex);
+ throw new KafkaStartingException("Error while creating Kafka topics", ex);
} catch (InterruptedException ex) {
Thread.currentThread().interrupt();
- throw new KafkaStartingException("Error when topic creating, ", ex);
+ throw new KafkaStartingException("Interrupted while creating Kafka topics", ex);
}
}
default void deleteTopics(List excludedTopics) {
try (var admin = createAdminClient()) {
var existingTopics = admin.listTopics().names().get(WAIT_TIMEOUT, TimeUnit.SECONDS);
- if (existingTopics.isEmpty()) {
- return;
- }
- var topics = topics().stream()
+ var topicsToDelete = topics().stream()
.filter(topic -> !excludedTopics.contains(topic))
- .toList();
- var topicsToDelete = topics.stream()
.filter(existingTopics::contains)
.toList();
if (topicsToDelete.isEmpty()) {
return;
}
- var topicsResult = admin.deleteTopics(topicsToDelete);
+
+ admin.deleteTopics(topicsToDelete).all().get(WAIT_TIMEOUT, TimeUnit.SECONDS);
Awaitility.await()
.atMost(Duration.ofSeconds(WAIT_TIMEOUT))
- .pollInterval(Duration.ofSeconds(2))
- .untilAsserted(() -> topicsResult.all().get(1, TimeUnit.SECONDS));
- var topicsAfterDelete = admin.listTopics().names().get(WAIT_TIMEOUT, TimeUnit.SECONDS);
- log.info("Topics list from 'AdminClient' after [TOPICS DELETED]: {}", topicsAfterDelete);
- assertThat(topicsAfterDelete.stream().noneMatch(topicsToDelete::contains))
- .isTrue();
- var actual = execInContainerKafkaTopicsListCommand();
- assertThat(topicsToDelete.stream().noneMatch(actual::contains))
- .isTrue();
+ .pollInterval(Duration.ofSeconds(1))
+ .until(() -> {
+ var actual = admin.listTopics().names().get(5, TimeUnit.SECONDS);
+ return actual.stream().noneMatch(topicsToDelete::contains);
+ });
+
+ var actual = admin.listTopics().names().get(WAIT_TIMEOUT, TimeUnit.SECONDS);
+ var remaining = topicsToDelete.stream().filter(actual::contains).toList();
+ if (!remaining.isEmpty()) {
+ throw new KafkaStartingException("Kafka topics were not deleted: " + remaining, null);
+ }
+ log.info("Kafka topics were deleted: {}", topicsToDelete);
} catch (ExecutionException | TimeoutException ex) {
- throw new KafkaStartingException("Error when topic deleting, ", ex);
+ throw new KafkaStartingException("Error while deleting Kafka topics", ex);
} catch (InterruptedException ex) {
Thread.currentThread().interrupt();
- throw new KafkaStartingException("Error when topic deleting, ", ex);
+ throw new KafkaStartingException("Interrupted while deleting Kafka topics", ex);
}
}
@@ -114,16 +116,29 @@ default AdminClient createAdminClient() {
}
default String execInContainerKafkaTopicsListCommandWithPath(String kafkaTopicsPath) {
- var kafkaTopicsListCommand = kafkaTopicsPath + " --bootstrap-server localhost:9093 --list";
+ var command = kafkaTopicsPath + " --bootstrap-server localhost:9093 --list";
try {
- var stdout = execInContainer("/bin/bash", "-c", kafkaTopicsListCommand).getStdout();
- log.info("Topics list from '{}': [{}]", kafkaTopicsPath, stdout.replace("\n", ","));
- return stdout;
+ var result = execInContainer("/bin/bash", "-c", command);
+ if (result.getExitCode() != 0) {
+ throw new KafkaStartingException(
+ "Kafka topics command failed with exit code " + result.getExitCode() +
+ ": " + result.getStderr(),
+ null);
+ }
+ log.info("Topics list from '{}': {}", kafkaTopicsPath, result.getStdout().replace("\n", ","));
+ return result.getStdout();
} catch (IOException ex) {
- throw new KafkaStartingException("Error when " + kafkaTopicsListCommand + ", ", ex);
+ throw new KafkaStartingException("Error while executing " + command, ex);
} catch (InterruptedException ex) {
Thread.currentThread().interrupt();
- throw new KafkaStartingException("Error when " + kafkaTopicsListCommand + ", ", ex);
+ throw new KafkaStartingException("Interrupted while executing " + command, ex);
}
}
+
+ default Set parseTopicList(String stdout) {
+ return Arrays.stream(stdout.split("\\R"))
+ .map(String::trim)
+ .filter(topic -> !topic.isEmpty())
+ .collect(java.util.stream.Collectors.toUnmodifiableSet());
+ }
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainer.java b/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainer.java
index c3fe90c1..90d79cb7 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainer.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainer.java
@@ -129,7 +129,7 @@
* Топики, которые не нужно очищать между тестами.
* Используется только если {@link #truncateTopics()} = true
*
- * пример — excludeTruncateTopics = {"kafka.topics.invoicing.id"}
+ * пример — excludeTruncateTopics = {"invoicing"} (реальное имя топика, не ключ property)
*/
String[] excludeTruncateTopics() default {};
diff --git a/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainerExtension.java b/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainerExtension.java
index d1eb2d37..5f8bcf48 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainerExtension.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainerExtension.java
@@ -1,9 +1,10 @@
package dev.vality.testcontainers.annotations.kafka;
+import dev.vality.testcontainers.annotations.kafka.constants.Provider;
import dev.vality.testcontainers.annotations.util.GenericContainerUtil;
+import dev.vality.testcontainers.annotations.util.SharedTestResourceLock;
import dev.vality.testcontainers.annotations.util.SpringApplicationPropertiesLoader;
-import lombok.extern.slf4j.Slf4j;
-import org.apache.kafka.clients.admin.AdminClient;
+import dev.vality.testcontainers.annotations.util.TestExecutionLock;
import org.junit.jupiter.api.extension.AfterAllCallback;
import org.junit.jupiter.api.extension.BeforeAllCallback;
import org.junit.jupiter.api.extension.BeforeEachCallback;
@@ -14,131 +15,164 @@
import org.springframework.test.context.ContextConfigurationAttributes;
import org.springframework.test.context.ContextCustomizer;
import org.springframework.test.context.ContextCustomizerFactory;
+import org.springframework.test.context.MergedContextConfiguration;
import java.util.Arrays;
import java.util.List;
import java.util.Optional;
-import java.util.stream.Collectors;
-
-/**
- * {@code @KafkaTestcontainerExtension} инициализирует тестконтейнер из {@link KafkaTestcontainerFactory},
- * настраивает, стартует, валидирует и останавливает
- *
{@link KafkaTestcontainerExtension.KafkaTestcontainerContextCustomizerFactory}
- * Инициализация настроек контейнеров в спринговый контекст тестового приложения реализован
- * под капотом аннотаций, на уровне реализации интерфейса —
- * информация о настройках используемого тестконтейнера и передаваемые через параметры аннотации настройки
- * инициализируются через {@link TestPropertyValues} и сливаются с текущим получаемым контекстом
- * приложения {@link ConfigurableApplicationContext}
- *
Инициализация кастомизированных фабрик с инициализацией настроек осуществляется через описание бинов
- * в файле META-INF/spring.factories
- *
Нюансы
- * Данное расширение немного сложнее других аналогичных в библиотеке за счет дополнительной работы с топиками
- *
Работа заключается в загрузке имен топиков из файла с настройками спринга {@link #loadTopics(String[])},
- * создании топиков через {@link AdminClient} в {@link KafkaContainerExtension#createTopics(List)},
- * а также валидации результата создания через запрос '/usr/bin/kafka-topics --zookeeper localhost:2181 --list'
- * напрямую в контейнере в {@link KafkaContainerExtension#execInContainerKafkaTopicsListCommand()}
- *
Также помимо перечисленного, при работе расширения для создания синглтона перед запуском тестов
- * в каждом файле будет проводится удаление созданных ранее топиков в {@link KafkaContainerExtension#deleteTopics(List)}
- * и дальнейшее пересоздание топиков в {@link KafkaContainerExtension#createTopics(List)},
- * таким образом обеспечивая изоляцию данных между файлами с тестами
- *
- * @see KafkaTestcontainerFactory KafkaTestcontainerFactory
- * @see KafkaTestcontainerExtension.KafkaTestcontainerContextCustomizerFactory KafkaTestcontainerContextCustomizerFactory
- * @see TestPropertyValues TestPropertyValues
- * @see ConfigurableApplicationContext ConfigurableApplicationContext
- * @see BeforeAllCallback BeforeAllCallback
- * @see AfterAllCallback AfterAllCallback
- */
-@Slf4j
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.ConcurrentMap;
+
public class KafkaTestcontainerExtension implements BeforeAllCallback, AfterAllCallback, BeforeEachCallback {
- private static final ThreadLocal THREAD_CONTAINER = new ThreadLocal<>();
+ private static final ConcurrentMap, ContainerReference> CONTAINERS = new ConcurrentHashMap<>();
@Override
public void beforeAll(ExtensionContext context) {
- if (findPrototypeAnnotation(context).isPresent()) {
- var annotation = findPrototypeAnnotation(context).get();
- var topics = loadTopics(annotation.topicsKeys());
- var container = KafkaTestcontainerFactory.container(annotation.provider(), topics);
- GenericContainerUtil.startContainer(container);
- container.createTopics(List.of());
- THREAD_CONTAINER.set(container);
- } else if (findSingletonAnnotation(context).isPresent()) {
- var annotation = findSingletonAnnotation(context).get();
- var topics = loadTopics(annotation.topicsKeys());
- var container = KafkaTestcontainerFactory.singletonContainer(annotation.provider(), topics);
- if (!container.isRunning()) {
- GenericContainerUtil.startContainer(container);
- container.createTopics(List.of());
- } else if (annotation.truncateTopics()) {
- var excludedTopics = Optional.ofNullable(annotation.excludeTruncateTopics())
- .map(List::of)
- .orElse(List.of());
- container.deleteTopics(excludedTopics);
- container.createTopics(excludedTopics);
- }
- THREAD_CONTAINER.set(container);
- }
+ var testClass = context.getRequiredTestClass();
+ getOrStart(testClass, configuration(testClass));
}
@Override
public void beforeEach(ExtensionContext context) {
- var container = THREAD_CONTAINER.get();
- if (findPrototypeAnnotation(context).isPresent()) {
- var annotation = findPrototypeAnnotation(context).get();
- if (container != null && container.isRunning() && annotation.truncateTopics()) {
- var excludedTopics = Optional.ofNullable(annotation.excludeTruncateTopics())
- .map(List::of)
- .orElse(List.of());
- container.deleteTopics(excludedTopics);
- container.createTopics(excludedTopics);
- }
- } else if (findSingletonAnnotation(context).isPresent()) {
- var annotation = findSingletonAnnotation(context).get();
- if (container != null && container.isRunning() && annotation.truncateTopics()) {
- var excludedTopics = Optional.ofNullable(annotation.excludeTruncateTopics())
- .map(List::of)
- .orElse(List.of());
- container.deleteTopics(excludedTopics);
- container.createTopics(excludedTopics);
+ var testClass = context.getRequiredTestClass();
+ var configuration = configuration(testClass);
+ if (configuration.truncateTopics()) {
+ TestExecutionLock.acquire(context);
+ try {
+ var reference = getOrStart(testClass, configuration);
+ reference.container().deleteTopics(configuration.excludedTopics());
+ reference.container().createTopics(List.of());
+ } catch (RuntimeException ex) {
+ TestExecutionLock.release(context);
+ throw ex;
}
}
}
@Override
public void afterAll(ExtensionContext context) {
- if (findPrototypeAnnotation(context).isPresent()) {
- var container = THREAD_CONTAINER.get();
- if (container != null && container.isRunning()) {
- container.stop();
+ var testClass = context.getRequiredTestClass();
+ var reference = CONTAINERS.remove(testClass);
+ try {
+ if (reference != null && !reference.singleton()) {
+ reference.container().stop();
+ }
+ } finally {
+ if (reference != null && reference.singleton()) {
+ SharedTestResourceLock.release(testClass);
}
- THREAD_CONTAINER.remove();
- } else if (findSingletonAnnotation(context).isPresent()) {
- THREAD_CONTAINER.remove();
}
}
- private static Optional findPrototypeAnnotation(ExtensionContext context) {
- return AnnotationSupport.findAnnotation(context.getTestClass(), KafkaTestcontainer.class);
+ private static ContainerReference getOrStart(Class> testClass, Configuration configuration) {
+ return CONTAINERS.computeIfAbsent(testClass, ignored -> createAndStart(testClass, configuration));
}
- private static Optional findPrototypeAnnotation(Class> testClass) {
- return AnnotationSupport.findAnnotation(testClass, KafkaTestcontainer.class);
+ private static ContainerReference createAndStart(Class> testClass, Configuration configuration) {
+ if (configuration.singleton()) {
+ SharedTestResourceLock.acquire(testClass);
+ }
+ try {
+ var container = configuration.singleton()
+ ? KafkaTestcontainerFactory.singletonContainer(configuration.provider(), configuration.topics())
+ : KafkaTestcontainerFactory.container(configuration.provider(), configuration.topics());
+ try {
+ GenericContainerUtil.startContainer(container);
+ if (configuration.singleton() && configuration.truncateTopics()) {
+ container.deleteTopics(configuration.excludedTopics());
+ }
+ container.createTopics(List.of());
+ return new ContainerReference(container, configuration.singleton());
+ } catch (RuntimeException ex) {
+ try {
+ if (configuration.singleton()) {
+ KafkaTestcontainerFactory.discardSingleton(container);
+ } else {
+ container.stop();
+ }
+ } catch (RuntimeException stopException) {
+ ex.addSuppressed(stopException);
+ }
+ throw ex;
+ }
+ } catch (RuntimeException ex) {
+ if (configuration.singleton()) {
+ SharedTestResourceLock.release(testClass);
+ }
+ throw ex;
+ }
}
- private static Optional findSingletonAnnotation(ExtensionContext context) {
- return AnnotationSupport.findAnnotation(context.getTestClass(), KafkaTestcontainerSingleton.class);
+ private static Configuration configuration(Class> testClass) {
+ var prototype = findPrototypeAnnotation(testClass);
+ if (prototype.isPresent()) {
+ var annotation = prototype.get();
+ return new Configuration(
+ false,
+ annotation.provider(),
+ loadTopics(annotation.topicsKeys()),
+ annotation.truncateTopics(),
+ List.copyOf(Arrays.asList(annotation.excludeTruncateTopics())),
+ List.copyOf(Arrays.asList(annotation.properties())));
+ }
+ var annotation = findSingletonAnnotation(testClass)
+ .orElseThrow(() -> new IllegalStateException("Kafka test annotation not found"));
+ return new Configuration(
+ true,
+ annotation.provider(),
+ loadTopics(annotation.topicsKeys()),
+ annotation.truncateTopics(),
+ List.copyOf(Arrays.asList(annotation.excludeTruncateTopics())),
+ List.copyOf(Arrays.asList(annotation.properties())));
+ }
+
+ private static Optional findPrototypeAnnotation(Class> testClass) {
+ return AnnotationSupport.findAnnotation(testClass, KafkaTestcontainer.class);
}
private static Optional findSingletonAnnotation(Class> testClass) {
return AnnotationSupport.findAnnotation(testClass, KafkaTestcontainerSingleton.class);
}
- private List loadTopics(String[] topicsKeys) {
- return SpringApplicationPropertiesLoader.loadFromSpringApplicationPropertiesFile(Arrays.asList(topicsKeys))
- .values().stream()
- .map(String::valueOf)
- .collect(Collectors.toList());
+ private static List loadTopics(String[] topicKeys) {
+ var properties = SpringApplicationPropertiesLoader.loadFromSpringApplicationPropertiesFile(
+ Arrays.asList(topicKeys));
+ return Arrays.stream(topicKeys)
+ .map(properties::getProperty)
+ .map(String::trim)
+ .filter(topic -> !topic.isEmpty())
+ .distinct()
+ .toList();
+ }
+
+ private static final class ContainerReference {
+
+ private final KafkaContainerExtension container;
+ private final boolean singleton;
+
+ private ContainerReference(KafkaContainerExtension container, boolean singleton) {
+ this.container = container;
+ this.singleton = singleton;
+ }
+
+ private KafkaContainerExtension container() {
+ return container;
+ }
+
+ private boolean singleton() {
+ return singleton;
+ }
+
+ }
+
+ private record Configuration(
+ boolean singleton,
+ Provider provider,
+ List topics,
+ boolean truncateTopics,
+ List excludedTopics,
+ List properties) {
}
public static class KafkaTestcontainerContextCustomizerFactory implements ContextCustomizerFactory {
@@ -147,22 +181,27 @@ public static class KafkaTestcontainerContextCustomizerFactory implements Contex
public ContextCustomizer createContextCustomizer(
Class> testClass,
List configAttributes) {
- return (context, mergedConfig) -> {
- if (findPrototypeAnnotation(testClass).isPresent()) {
- init(context, findPrototypeAnnotation(testClass).get().properties());
- } else if (findSingletonAnnotation(testClass).isPresent()) {
- init(context, findSingletonAnnotation(testClass).get().properties());
- }
- };
+ if (findPrototypeAnnotation(testClass).isEmpty() && findSingletonAnnotation(testClass).isEmpty()) {
+ return null;
+ }
+ return new KafkaContextCustomizer(testClass, configuration(testClass));
}
+ }
- private void init(ConfigurableApplicationContext context, String[] properties) {
- var container = THREAD_CONTAINER.get();
+ private record KafkaContextCustomizer(
+ Class> testClass,
+ Configuration configuration) implements ContextCustomizer {
+
+ @Override
+ public void customizeContext(
+ ConfigurableApplicationContext context,
+ MergedContextConfiguration mergedConfig) {
+ var container = getOrStart(testClass, configuration).container();
TestPropertyValues.of(
"kafka.bootstrap-servers=" + container.getBootstrapServers(),
"spring.kafka.bootstrap-servers=" + container.getBootstrapServers(),
"kafka.ssl.enabled=false")
- .and(properties)
+ .and(configuration.properties().toArray(String[]::new))
.applyTo(context);
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainerFactory.java b/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainerFactory.java
index 5c3aa57b..eabd5dd9 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainerFactory.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainerFactory.java
@@ -1,57 +1,71 @@
package dev.vality.testcontainers.annotations.kafka;
import dev.vality.testcontainers.annotations.kafka.constants.Provider;
+import dev.vality.testcontainers.annotations.util.ContainerShutdownRegistry;
import lombok.AccessLevel;
import lombok.NoArgsConstructor;
-import lombok.Synchronized;
-import lombok.extern.slf4j.Slf4j;
import java.util.List;
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.ConcurrentMap;
/**
- * Фабрика по созданию контейнеров
- * {@link #create(Provider, List)} создает экземпляр тестконтейнера
- *
{@link #getOrCreateSingletonContainer(Provider, List)} создает синглтон тестконтейнера
- *
- * @see KafkaTestcontainerExtension KafkaTestcontainerExtension
+ * Фабрика по созданию Kafka контейнеров.
*/
-@Slf4j
@NoArgsConstructor(access = AccessLevel.PRIVATE)
public class KafkaTestcontainerFactory {
- private KafkaContainerExtension kafkaContainer;
+ private final ConcurrentMap singletonContainers = new ConcurrentHashMap<>();
public static KafkaContainerExtension container(Provider provider, List topics) {
- return instance().create(provider, topics);
+ return instance().create(Config.of(provider, topics));
}
public static KafkaContainerExtension singletonContainer(Provider provider, List topics) {
- return instance().getOrCreateSingletonContainer(provider, topics);
+ var config = Config.of(provider, topics);
+ return instance().singletonContainers.computeIfAbsent(
+ config,
+ key -> ContainerShutdownRegistry.register(instance().create(key)));
+ }
+
+ static void discardSingleton(KafkaContainerExtension container) {
+ if (instance().singletonContainers.values().remove(container)) {
+ ContainerShutdownRegistry.unregister(container);
+ container.stop();
+ }
}
private static KafkaTestcontainerFactory instance() {
return SingletonHolder.INSTANCE;
}
- @Synchronized
- private KafkaContainerExtension getOrCreateSingletonContainer(Provider provider, List topics) {
- if (kafkaContainer != null) {
- return kafkaContainer;
- }
- kafkaContainer = create(provider, topics);
- return kafkaContainer;
+ private KafkaContainerExtension create(Config config) {
+ return switch (config.provider()) {
+ case APACHE -> new ApacheKafkaContainer(config.topics());
+ case CONFLUENT -> new ConfluentKafkaContainer(config.topics());
+ };
}
- private KafkaContainerExtension create(Provider provider, List topics) {
- return switch (provider) {
- case APACHE -> new ApacheKafkaContainer(topics);
- case CONFLUENT -> new ConfluentKafkaContainer(topics);
- };
+ private record Config(Provider provider, List topics) {
+
+ private static Config of(Provider provider, List topics) {
+ if (provider == null) {
+ throw new IllegalArgumentException("Kafka provider must not be null");
+ }
+ var canonicalTopics = topics == null
+ ? List.of()
+ : topics.stream()
+ .map(String::trim)
+ .filter(topic -> !topic.isEmpty())
+ .distinct()
+ .sorted()
+ .toList();
+ return new Config(provider, List.copyOf(canonicalTopics));
+ }
}
private static class SingletonHolder {
private static final KafkaTestcontainerFactory INSTANCE = new KafkaTestcontainerFactory();
-
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainerSingleton.java b/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainerSingleton.java
index 64631080..5d819d62 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainerSingleton.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/kafka/KafkaTestcontainerSingleton.java
@@ -142,7 +142,7 @@
* Топики, которые не нужно очищать между тестами.
* Используется только если {@link #truncateTopics()} = true
*
- * пример — excludeTruncateTopics = {"kafka.topics.invoicing.id"}
+ * пример — excludeTruncateTopics = {"invoicing"} (реальное имя топика, не ключ property)
*/
String[] excludeTruncateTopics() default {};
diff --git a/src/main/java/dev/vality/testcontainers/annotations/kafka/config/KafkaConsumer.java b/src/main/java/dev/vality/testcontainers/annotations/kafka/config/KafkaConsumer.java
index 8383add9..6074ca83 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/kafka/config/KafkaConsumer.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/kafka/config/KafkaConsumer.java
@@ -1,11 +1,11 @@
package dev.vality.testcontainers.annotations.kafka.config;
import dev.vality.kafka.common.serialization.AbstractThriftDeserializer;
-import dev.vality.testcontainers.annotations.kafka.KafkaTestcontainer;
import lombok.RequiredArgsConstructor;
import org.apache.kafka.clients.consumer.ConsumerConfig;
import org.apache.kafka.common.serialization.StringDeserializer;
import org.apache.thrift.TBase;
+import org.springframework.beans.factory.DisposableBean;
import org.springframework.kafka.core.DefaultKafkaConsumerFactory;
import org.springframework.kafka.listener.ConcurrentMessageListenerContainer;
import org.springframework.kafka.listener.ContainerProperties;
@@ -13,40 +13,80 @@
import java.util.HashMap;
import java.util.Map;
+import java.util.Set;
import java.util.UUID;
+import java.util.concurrent.ConcurrentHashMap;
/**
- * Листенер для чтения данных из тестового трифтового топика
- * Для получения конкретного сообщения необходимо имплементировать в тесте интерфейс
- * {@link MessageListener}
- *
Пример использования {@link KafkaTestcontainer} с {@link KafkaConsumer} — в
- * sink-drinker
- *
Пример
- * {@code
- * @Autowired
- * private KafkaConsumer testPayoutEventKafkaConsumer;
- *
- * ...
- *
- * testPayoutEventKafkaConsumer.read(topicName, data -> readEvents.add(data.value()));
- * Unreliables.retryUntilTrue(TIMEOUT, TimeUnit.SECONDS, () -> readEvents.size() == expected);
- *
- * ...
- * }
- *
- * @see KafkaConsumerConfig KafkaConsumerConfig
+ * Listener helper for reading records from a test Thrift topic.
*/
@RequiredArgsConstructor
-public class KafkaConsumer> {
+public class KafkaConsumer> implements DisposableBean, AutoCloseable {
private final String bootstrapAddress;
private final AbstractThriftDeserializer deserializer;
+ private final Set> containers = ConcurrentHashMap.newKeySet();
+ /**
+ * Starts a listener container managed by this bean.
+ */
public void read(String topic, MessageListener messageListener) {
+ start(topic, messageListener);
+ }
+
+ /**
+ * Starts a listener container and returns a handle that can be stopped explicitly.
+ */
+ public ConcurrentMessageListenerContainer start(
+ String topic,
+ MessageListener messageListener) {
var container = new ConcurrentMessageListenerContainer<>(
consumerFactory(),
containerProperties(topic, messageListener));
- container.start();
+ containers.add(container);
+ try {
+ container.start();
+ return container;
+ } catch (RuntimeException ex) {
+ try {
+ container.stop();
+ containers.remove(container);
+ } catch (RuntimeException stopException) {
+ ex.addSuppressed(stopException);
+ }
+ throw ex;
+ }
+ }
+
+ public void stop(ConcurrentMessageListenerContainer container) {
+ if (container != null && containers.contains(container)) {
+ container.stop();
+ containers.remove(container);
+ }
+ }
+
+ @Override
+ public void destroy() {
+ close();
+ }
+
+ @Override
+ public void close() {
+ RuntimeException failure = null;
+ for (var container : Set.copyOf(containers)) {
+ try {
+ container.stop();
+ containers.remove(container);
+ } catch (RuntimeException ex) {
+ if (failure == null) {
+ failure = new IllegalStateException("Unable to stop Kafka listener containers");
+ }
+ failure.addSuppressed(ex);
+ }
+ }
+ if (failure != null) {
+ throw failure;
+ }
}
private ContainerProperties containerProperties(String topic, MessageListener messageListener) {
@@ -64,6 +104,7 @@ private Map consumerConfig() {
properties.put(ConsumerConfig.BOOTSTRAP_SERVERS_CONFIG, bootstrapAddress);
properties.put(ConsumerConfig.GROUP_ID_CONFIG, UUID.randomUUID().toString());
properties.put(ConsumerConfig.AUTO_OFFSET_RESET_CONFIG, "earliest");
+ properties.put(ConsumerConfig.ENABLE_AUTO_COMMIT_CONFIG, false);
return properties;
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/kafka/config/KafkaProducer.java b/src/main/java/dev/vality/testcontainers/annotations/kafka/config/KafkaProducer.java
index 80b9760e..eb15143b 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/kafka/config/KafkaProducer.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/kafka/config/KafkaProducer.java
@@ -1,27 +1,11 @@
package dev.vality.testcontainers.annotations.kafka.config;
-import dev.vality.testcontainers.annotations.kafka.KafkaTestcontainer;
import lombok.RequiredArgsConstructor;
import lombok.extern.slf4j.Slf4j;
import org.springframework.kafka.core.KafkaTemplate;
/**
- * Обертка над {@link KafkaTemplate}, используется для отправки сообщений в тестовый топик
- * Пример использования {@link KafkaTestcontainer} с {@link KafkaProducer} — в
- * magista
- *
Пример
- * {@code
- * @Autowired
- * private KafkaProducer> testThriftKafkaProducer;
- *
- * ...
- *
- * testThriftKafkaProducer.send(invoicingTopicName, sinkEvent);
- *
- * ...
- * }
- *
- * @see KafkaProducerTestConfig KafkaProducerTestConfig
+ * Обертка над {@link KafkaTemplate}, используется для отправки сообщений в тестовый топик.
*/
@RequiredArgsConstructor
@Slf4j
@@ -31,16 +15,11 @@ public class KafkaProducer {
public void send(String topic, T payload) {
log.info("Sending payload='{}' to topic='{}'", payload, topic);
- kafkaTemplate.send(topic, payload)
- .join();
- kafkaTemplate.getProducerFactory().reset();
+ kafkaTemplate.send(topic, payload).join();
}
public void send(String topic, String key, T payload) {
log.info("Sending key='{}' payload='{}' to topic='{}'", key, payload, topic);
- kafkaTemplate.send(topic, key, payload)
- .join();
- kafkaTemplate.getProducerFactory().reset();
+ kafkaTemplate.send(topic, key, payload).join();
}
-
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/kafka/config/KafkaProducerTestConfig.java b/src/main/java/dev/vality/testcontainers/annotations/kafka/config/KafkaProducerTestConfig.java
index 5bc32f72..4e709818 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/kafka/config/KafkaProducerTestConfig.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/kafka/config/KafkaProducerTestConfig.java
@@ -4,45 +4,60 @@
import org.apache.kafka.clients.producer.ProducerConfig;
import org.apache.kafka.common.serialization.StringSerializer;
import org.apache.thrift.TBase;
+import org.springframework.beans.factory.annotation.Qualifier;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.boot.test.context.TestConfiguration;
import org.springframework.context.annotation.Bean;
-import org.springframework.context.annotation.Primary;
import org.springframework.kafka.core.DefaultKafkaProducerFactory;
import org.springframework.kafka.core.KafkaTemplate;
-import org.springframework.kafka.core.ProducerFactory;
-import org.springframework.util.ObjectUtils;
+import org.springframework.util.StringUtils;
import java.util.HashMap;
import java.util.Map;
import java.util.UUID;
/**
- * Конфиг для инциализации тестового продьюссера для тестирования трифтовых топиков
- *
- * @see KafkaProducer KafkaProducer
+ * Конфигурация тестового producer для Thrift сообщений.
*/
@TestConfiguration("kafkaProducerTestContainersAnnotationsConfig")
public class KafkaProducerTestConfig {
- @Bean
- @Primary
- public String bootstrapAddress(@Value("${spring.kafka.bootstrap-servers:}") String primaryLocation,
- @Value("${kafka.bootstrap-servers:}") String secondaryLocation) {
- return !ObjectUtils.isEmpty(primaryLocation) ? primaryLocation : secondaryLocation;
+ public static final String BOOTSTRAP_ADDRESS_BEAN = "testKafkaBootstrapAddress";
+ public static final String PRODUCER_FACTORY_BEAN = "testThriftKafkaProducerFactory";
+ public static final String KAFKA_TEMPLATE_BEAN = "testThriftKafkaTemplate";
+
+ @Bean(BOOTSTRAP_ADDRESS_BEAN)
+ public String bootstrapAddress(
+ @Value("${spring.kafka.bootstrap-servers:}") String primaryLocation,
+ @Value("${kafka.bootstrap-servers:}") String secondaryLocation) {
+ var address = StringUtils.hasText(primaryLocation) ? primaryLocation : secondaryLocation;
+ if (!StringUtils.hasText(address)) {
+ throw new IllegalStateException("Kafka bootstrap servers are not configured");
+ }
+ return address;
}
- @Bean
- public KafkaProducer> testThriftKafkaProducer(String bootstrapAddress) {
- return new KafkaProducer<>(new KafkaTemplate<>(thriftProducerFactory(bootstrapAddress)));
+ @Bean(value = PRODUCER_FACTORY_BEAN, destroyMethod = "destroy")
+ public DefaultKafkaProducerFactory> thriftProducerFactory(
+ @Qualifier(BOOTSTRAP_ADDRESS_BEAN) String bootstrapAddress) {
+ Map properties = new HashMap<>();
+ properties.put(ProducerConfig.BOOTSTRAP_SERVERS_CONFIG, bootstrapAddress);
+ properties.put(ProducerConfig.CLIENT_ID_CONFIG, UUID.randomUUID().toString());
+ properties.put(ProducerConfig.KEY_SERIALIZER_CLASS_CONFIG, StringSerializer.class.getName());
+ properties.put(ProducerConfig.VALUE_SERIALIZER_CLASS_CONFIG, ThriftSerializer.class.getName());
+ return new DefaultKafkaProducerFactory<>(properties);
}
- private ProducerFactory> thriftProducerFactory(String bootstrapAddress) {
- Map props = new HashMap<>();
- props.put(ProducerConfig.BOOTSTRAP_SERVERS_CONFIG, bootstrapAddress);
- props.put(ProducerConfig.CLIENT_ID_CONFIG, UUID.randomUUID().toString());
- props.put(ProducerConfig.KEY_SERIALIZER_CLASS_CONFIG, StringSerializer.class.getName());
- props.put(ProducerConfig.VALUE_SERIALIZER_CLASS_CONFIG, ThriftSerializer.class.getName());
- return new DefaultKafkaProducerFactory<>(props);
+ @Bean(KAFKA_TEMPLATE_BEAN)
+ public KafkaTemplate> thriftKafkaTemplate(
+ @Qualifier(PRODUCER_FACTORY_BEAN)
+ DefaultKafkaProducerFactory> producerFactory) {
+ return new KafkaTemplate<>(producerFactory);
+ }
+
+ @Bean
+ public KafkaProducer> testThriftKafkaProducer(
+ @Qualifier(KAFKA_TEMPLATE_BEAN) KafkaTemplate> kafkaTemplate) {
+ return new KafkaProducer<>(kafkaTemplate);
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/minio/MinioBucketManager.java b/src/main/java/dev/vality/testcontainers/annotations/minio/MinioBucketManager.java
new file mode 100644
index 00000000..d535efbd
--- /dev/null
+++ b/src/main/java/dev/vality/testcontainers/annotations/minio/MinioBucketManager.java
@@ -0,0 +1,58 @@
+package dev.vality.testcontainers.annotations.minio;
+
+import io.minio.*;
+import lombok.AccessLevel;
+import lombok.NoArgsConstructor;
+import org.testcontainers.containers.GenericContainer;
+
+import java.util.List;
+
+import static dev.vality.testcontainers.annotations.util.SpringApplicationPropertiesLoader.loadDefaultLibraryProperty;
+
+@NoArgsConstructor(access = AccessLevel.PRIVATE)
+final class MinioBucketManager {
+
+ static void ensureBucket(GenericContainer> container, String bucketName) {
+ try (var client = client(container)) {
+ var exists = client.bucketExists(BucketExistsArgs.builder().bucket(bucketName).build());
+ if (!exists) {
+ client.makeBucket(MakeBucketArgs.builder().bucket(bucketName).build());
+ }
+ } catch (Exception ex) {
+ throw new IllegalStateException("Unable to create MinIO bucket '" + bucketName + "'", ex);
+ }
+ }
+
+ static void cleanupBucket(
+ GenericContainer> container,
+ String bucketName,
+ List excludedPrefixes) {
+ try (var client = client(container)) {
+ var objects = client.listObjects(ListObjectsArgs.builder()
+ .bucket(bucketName)
+ .recursive(true)
+ .build());
+ for (var objectResult : objects) {
+ var objectName = objectResult.get().objectName();
+ if (excludedPrefixes.stream().noneMatch(objectName::startsWith)) {
+ client.removeObject(RemoveObjectArgs.builder()
+ .bucket(bucketName)
+ .object(objectName)
+ .build());
+ }
+ }
+ } catch (Exception ex) {
+ throw new IllegalStateException("Unable to clean MinIO bucket '" + bucketName + "'", ex);
+ }
+ }
+
+ private static MinioClient client(GenericContainer> container) {
+ var endpoint = "http://%s:%d".formatted(container.getHost(), container.getMappedPort(9000));
+ return MinioClient.builder()
+ .endpoint(endpoint)
+ .credentials(
+ loadDefaultLibraryProperty(MinioTestcontainerFactory.MINIO_USER),
+ loadDefaultLibraryProperty(MinioTestcontainerFactory.MINIO_PASSWORD))
+ .build();
+ }
+}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainer.java b/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainer.java
index 96064414..9c6fcdc1 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainer.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainer.java
@@ -48,4 +48,14 @@
*/
String bucketName() default "test";
+ /**
+ * Очищать содержимое bucket перед каждым тестом.
+ */
+ boolean cleanupBucket() default true;
+
+ /**
+ * Префиксы объектов, которые не нужно удалять при очистке bucket.
+ */
+ String[] excludeCleanupPrefixes() default {};
+
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainerExtension.java b/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainerExtension.java
index 77231a44..3ae8d997 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainerExtension.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainerExtension.java
@@ -1,9 +1,11 @@
package dev.vality.testcontainers.annotations.minio;
import dev.vality.testcontainers.annotations.util.GenericContainerUtil;
-import lombok.extern.slf4j.Slf4j;
+import dev.vality.testcontainers.annotations.util.SharedTestResourceLock;
+import dev.vality.testcontainers.annotations.util.TestExecutionLock;
import org.junit.jupiter.api.extension.AfterAllCallback;
import org.junit.jupiter.api.extension.BeforeAllCallback;
+import org.junit.jupiter.api.extension.BeforeEachCallback;
import org.junit.jupiter.api.extension.ExtensionContext;
import org.junit.platform.commons.support.AnnotationSupport;
import org.springframework.boot.test.util.TestPropertyValues;
@@ -11,136 +13,201 @@
import org.springframework.test.context.ContextConfigurationAttributes;
import org.springframework.test.context.ContextCustomizer;
import org.springframework.test.context.ContextCustomizerFactory;
+import org.springframework.test.context.MergedContextConfiguration;
import org.testcontainers.containers.GenericContainer;
+import java.util.Arrays;
import java.util.List;
import java.util.Optional;
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.ConcurrentMap;
import static dev.vality.testcontainers.annotations.util.SpringApplicationPropertiesLoader.loadDefaultLibraryProperty;
-/**
- * {@code @MinioTestcontainerExtension} инициализирует тестконтейнер из {@link MinioTestcontainerFactory},
- * настраивает, стартует, валидирует и останавливает
- *
{@link MinioTestcontainerExtension.MinioTestcontainerContextCustomizerFactory}
- * Инициализация настроек контейнеров в спринговый контекст тестового приложения реализован
- * под капотом аннотаций, на уровне реализации интерфейса —
- * информация о настройках используемого тестконтейнера и передаваемые через параметры аннотации настройки
- * инициализируются через {@link TestPropertyValues} и сливаются с текущим получаемым контекстом
- * приложения {@link ConfigurableApplicationContext}
- *
Инициализация кастомизированных фабрик с инициализацией настроек осуществляется через описание бинов
- * в файле META-INF/spring.factories
- *
- * @see MinioTestcontainerFactory MinioTestcontainerFactory
- * @see MinioTestcontainerExtension.MinioTestcontainerContextCustomizerFactory MinioTestcontainerContextCustomizerFactory
- * @see TestPropertyValues TestPropertyValues
- * @see ConfigurableApplicationContext ConfigurableApplicationContext
- * @see BeforeAllCallback BeforeAllCallback
- * @see AfterAllCallback AfterAllCallback
- */
-@Slf4j
-public class MinioTestcontainerExtension implements BeforeAllCallback, AfterAllCallback {
-
- private static final ThreadLocal> THREAD_CONTAINER = new ThreadLocal<>();
+public class MinioTestcontainerExtension implements BeforeAllCallback, BeforeEachCallback, AfterAllCallback {
+
+ private static final ConcurrentMap, ContainerReference> CONTAINERS = new ConcurrentHashMap<>();
@Override
public void beforeAll(ExtensionContext context) {
- if (findPrototypeAnnotation(context).isPresent()) {
- var container = MinioTestcontainerFactory.container();
- GenericContainerUtil.startContainer(container);
- THREAD_CONTAINER.set(container);
- } else if (findSingletonAnnotation(context).isPresent()) {
- var container = MinioTestcontainerFactory.singletonContainer();
- if (!container.isRunning()) {
- GenericContainerUtil.startContainer(container);
+ var testClass = context.getRequiredTestClass();
+ getOrStart(testClass, configuration(testClass));
+ }
+
+ @Override
+ public void beforeEach(ExtensionContext context) {
+ var testClass = context.getRequiredTestClass();
+ var configuration = configuration(testClass);
+ if (configuration.cleanupBucket()) {
+ TestExecutionLock.acquire(context);
+ try {
+ var reference = getOrStart(testClass, configuration);
+ MinioBucketManager.cleanupBucket(
+ reference.container(),
+ configuration.bucketName(),
+ configuration.excludedPrefixes());
+ } catch (RuntimeException ex) {
+ TestExecutionLock.release(context);
+ throw ex;
}
- THREAD_CONTAINER.set(container);
}
}
@Override
public void afterAll(ExtensionContext context) {
- if (findPrototypeAnnotation(context).isPresent()) {
- var container = THREAD_CONTAINER.get();
- if (container != null && container.isRunning()) {
- container.stop();
+ var testClass = context.getRequiredTestClass();
+ var reference = CONTAINERS.remove(testClass);
+ try {
+ if (reference != null && !reference.singleton()) {
+ reference.container().stop();
+ }
+ } finally {
+ if (reference != null && reference.singleton()) {
+ SharedTestResourceLock.release(testClass);
}
- THREAD_CONTAINER.remove();
- } else if (findSingletonAnnotation(context).isPresent()) {
- THREAD_CONTAINER.remove();
}
}
- private static Optional findPrototypeAnnotation(ExtensionContext context) {
- return AnnotationSupport.findAnnotation(context.getTestClass(), MinioTestcontainer.class);
+ private static ContainerReference getOrStart(Class> testClass, Configuration configuration) {
+ return CONTAINERS.computeIfAbsent(testClass, ignored -> createAndStart(testClass, configuration));
}
- private static Optional findPrototypeAnnotation(Class> testClass) {
- return AnnotationSupport.findAnnotation(testClass, MinioTestcontainer.class);
+ private static ContainerReference createAndStart(Class> testClass, Configuration configuration) {
+ if (configuration.singleton()) {
+ SharedTestResourceLock.acquire(testClass);
+ }
+ try {
+ var container = configuration.singleton()
+ ? MinioTestcontainerFactory.singletonContainer()
+ : MinioTestcontainerFactory.container();
+ try {
+ GenericContainerUtil.startContainer(container);
+ MinioBucketManager.ensureBucket(container, configuration.bucketName());
+ if (configuration.singleton() && configuration.cleanupBucket()) {
+ MinioBucketManager.cleanupBucket(
+ container,
+ configuration.bucketName(),
+ configuration.excludedPrefixes());
+ }
+ return new ContainerReference(container, configuration.singleton());
+ } catch (RuntimeException ex) {
+ try {
+ if (configuration.singleton()) {
+ MinioTestcontainerFactory.discardSingleton(container);
+ } else {
+ container.stop();
+ }
+ } catch (RuntimeException stopException) {
+ ex.addSuppressed(stopException);
+ }
+ throw ex;
+ }
+ } catch (RuntimeException ex) {
+ if (configuration.singleton()) {
+ SharedTestResourceLock.release(testClass);
+ }
+ throw ex;
+ }
}
- private static Optional findSingletonAnnotation(ExtensionContext context) {
- return AnnotationSupport.findAnnotation(context.getTestClass(), MinioTestcontainerSingleton.class);
+ private static Configuration configuration(Class> testClass) {
+ var prototype = findPrototypeAnnotation(testClass);
+ if (prototype.isPresent()) {
+ var annotation = prototype.get();
+ return new Configuration(
+ false,
+ annotation.bucketName(),
+ annotation.cleanupBucket(),
+ List.copyOf(Arrays.asList(annotation.excludeCleanupPrefixes())),
+ List.copyOf(Arrays.asList(annotation.properties())));
+ }
+ var annotation = findSingletonAnnotation(testClass)
+ .orElseThrow(() -> new IllegalStateException("MinIO test annotation not found"));
+ return new Configuration(
+ true,
+ annotation.bucketName(),
+ annotation.cleanupBucket(),
+ List.copyOf(Arrays.asList(annotation.excludeCleanupPrefixes())),
+ List.copyOf(Arrays.asList(annotation.properties())));
+ }
+
+ private static Optional findPrototypeAnnotation(Class> testClass) {
+ return AnnotationSupport.findAnnotation(testClass, MinioTestcontainer.class);
}
private static Optional findSingletonAnnotation(Class> testClass) {
return AnnotationSupport.findAnnotation(testClass, MinioTestcontainerSingleton.class);
}
+ private static final class ContainerReference {
+
+ private final GenericContainer> container;
+ private final boolean singleton;
+
+ private ContainerReference(GenericContainer> container, boolean singleton) {
+ this.container = container;
+ this.singleton = singleton;
+ }
+
+ private GenericContainer> container() {
+ return container;
+ }
+
+ private boolean singleton() {
+ return singleton;
+ }
+
+ }
+
+ private record Configuration(
+ boolean singleton,
+ String bucketName,
+ boolean cleanupBucket,
+ List excludedPrefixes,
+ List properties) {
+ }
+
public static class MinioTestcontainerContextCustomizerFactory implements ContextCustomizerFactory {
@Override
public ContextCustomizer createContextCustomizer(
Class> testClass,
List configAttributes) {
- return (context, mergedConfig) -> {
- if (findPrototypeAnnotation(testClass).isPresent()) {
- var annotation = findPrototypeAnnotation(testClass).get();
- init(
- context,
- annotation.bucketName(),
- annotation.properties());
- } else {
- findSingletonAnnotation(testClass).ifPresent(
- annotation -> init(
- context,
- annotation.bucketName(),
- annotation.properties()));
- }
- };
+ if (findPrototypeAnnotation(testClass).isEmpty() && findSingletonAnnotation(testClass).isEmpty()) {
+ return null;
+ }
+ return new MinioContextCustomizer(testClass, configuration(testClass));
}
+ }
- private void init(
+ private record MinioContextCustomizer(
+ Class> testClass,
+ Configuration configuration) implements ContextCustomizer {
+
+ @Override
+ public void customizeContext(
ConfigurableApplicationContext context,
- String bucketName,
- String[] properties) {
- var container = THREAD_CONTAINER.get();
+ MergedContextConfiguration mergedConfig) {
+ var container = getOrStart(testClass, configuration).container();
+ var hostAndPort = container.getHost() + ":" + container.getMappedPort(9000);
+ var endpoint = "http://" + hostAndPort + "/";
+ var user = loadDefaultLibraryProperty(MinioTestcontainerFactory.MINIO_USER);
+ var password = loadDefaultLibraryProperty(MinioTestcontainerFactory.MINIO_PASSWORD);
TestPropertyValues.of(
- // deprecated
- "storage.endpoint=" + container.getHost() + ":" +
- container.getMappedPort(9000),
-// "storage.signingRegion=" + signingRegion,
- "storage.accessKey=" + loadDefaultLibraryProperty(MinioTestcontainerFactory.MINIO_USER),
- "storage.secretKey=" + loadDefaultLibraryProperty(MinioTestcontainerFactory.MINIO_PASSWORD),
-// "storage.clientProtocol=" + clientProtocol,
-// "storage.clientMaxErrorRetry=" + clientMaxErrorRetry,
- "storage.bucketName=" + bucketName,
- // --
- "s3.endpoint=" + container.getHost() + ":" + container.getMappedPort(9000),
- "s3.bucket-name=" + bucketName,
-// "s3.signing-region=" + signingRegion,
-// "s3.client-protocol=" + clientProtocol,
-// "s3.client-max-error-retry=" + clientMaxErrorRetry,
-// "s3.signer-override=" + signerOverride,
- "s3.access-key=" + loadDefaultLibraryProperty(MinioTestcontainerFactory.MINIO_USER),
- "s3.secret-key=" + loadDefaultLibraryProperty(MinioTestcontainerFactory.MINIO_PASSWORD),
+ "storage.endpoint=" + hostAndPort,
+ "storage.accessKey=" + user,
+ "storage.secretKey=" + password,
+ "storage.bucketName=" + configuration.bucketName(),
+ "s3.endpoint=" + hostAndPort,
+ "s3.bucket-name=" + configuration.bucketName(),
+ "s3.access-key=" + user,
+ "s3.secret-key=" + password,
"s3-sdk-v2.enabled=false",
- "s3-sdk-v2.endpoint=" + String.format("http://%s:%d/", container.getHost(),
- container.getMappedPort(9000)),
- "s3-sdk-v2.bucket-name=" + bucketName,
-// "s3-sdk-v2.region=" + signingRegion,
- "s3-sdk-v2.access-key=" + loadDefaultLibraryProperty(MinioTestcontainerFactory.MINIO_USER),
- "s3-sdk-v2.secret-key=" + loadDefaultLibraryProperty(MinioTestcontainerFactory.MINIO_PASSWORD))
- .and(properties)
+ "s3-sdk-v2.endpoint=" + endpoint,
+ "s3-sdk-v2.bucket-name=" + configuration.bucketName(),
+ "s3-sdk-v2.access-key=" + user,
+ "s3-sdk-v2.secret-key=" + password)
+ .and(configuration.properties().toArray(String[]::new))
.applyTo(context);
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainerFactory.java b/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainerFactory.java
index 9945356b..6b3c2677 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainerFactory.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainerFactory.java
@@ -1,26 +1,18 @@
package dev.vality.testcontainers.annotations.minio;
+import dev.vality.testcontainers.annotations.util.ContainerShutdownRegistry;
import lombok.AccessLevel;
import lombok.NoArgsConstructor;
-import lombok.Synchronized;
import org.testcontainers.containers.GenericContainer;
-import org.testcontainers.containers.Network;
import org.testcontainers.utility.DockerImageName;
import java.time.Duration;
-import java.util.UUID;
import static dev.vality.testcontainers.annotations.util.GenericContainerUtil.getWaitStrategy;
import static dev.vality.testcontainers.annotations.util.SpringApplicationPropertiesLoader.loadDefaultLibraryProperty;
/**
- * Фабрика по созданию контейнеров
- * {@link #create()} создает экземпляр тестконтейнера
- *
{@link #getOrCreateSingletonContainer()} создает синглтон тестконтейнера
- *
{@link #MINIO_USER} необходимо указать в файле application.yml при необходимости другого ключа
- *
{@link #MINIO_PASSWORD} необходимо указать в файле application.yml при необходимости другого ключа
- *
- * @see MinioTestcontainerExtension MinioTestcontainerExtension
+ * Фабрика по созданию MinIO контейнеров.
*/
@NoArgsConstructor(access = AccessLevel.PRIVATE)
public class MinioTestcontainerFactory {
@@ -30,7 +22,7 @@ public class MinioTestcontainerFactory {
private static final String MINIO_IMAGE_NAME = "quay.io/minio/minio";
private static final String TAG_PROPERTY = "testcontainers.minio.tag";
- private GenericContainer> minioContainer;
+ private volatile GenericContainer> singletonContainer;
public static GenericContainer> container() {
return instance().create();
@@ -40,37 +32,51 @@ public static GenericContainer> singletonContainer() {
return instance().getOrCreateSingletonContainer();
}
+ static void discardSingleton(GenericContainer> container) {
+ instance().discard(container);
+ }
+
private static MinioTestcontainerFactory instance() {
return SingletonHolder.INSTANCE;
}
- @Synchronized
+ private void discard(GenericContainer> container) {
+ synchronized (this) {
+ if (singletonContainer != container) {
+ return;
+ }
+ singletonContainer = null;
+ }
+ ContainerShutdownRegistry.unregister(container);
+ container.stop();
+ }
+
private GenericContainer> getOrCreateSingletonContainer() {
- if (minioContainer != null) {
- return minioContainer;
+ var result = singletonContainer;
+ if (result == null) {
+ synchronized (this) {
+ result = singletonContainer;
+ if (result == null) {
+ result = ContainerShutdownRegistry.register(create());
+ singletonContainer = result;
+ }
+ }
}
- minioContainer = create();
- return minioContainer;
+ return result;
}
private GenericContainer> create() {
- GenericContainer> container = new GenericContainer<>(
- DockerImageName
- .parse(MINIO_IMAGE_NAME)
- .withTag(loadDefaultLibraryProperty(TAG_PROPERTY)))
+ return new GenericContainer<>(DockerImageName.parse(MINIO_IMAGE_NAME)
+ .withTag(loadDefaultLibraryProperty(TAG_PROPERTY)))
.withExposedPorts(9000)
.withEnv("MINIO_ROOT_USER", loadDefaultLibraryProperty(MINIO_USER))
.withEnv("MINIO_ROOT_PASSWORD", loadDefaultLibraryProperty(MINIO_PASSWORD))
- .withCommand("server /data{1...12}")
+ .withCommand("server /data")
.waitingFor(getWaitStrategy("/minio/health/live", 200, 9000, Duration.ofMinutes(1)));
- container.withNetworkAliases("minio-" + UUID.randomUUID());
- container.withNetwork(Network.SHARED);
- return container;
}
private static class SingletonHolder {
private static final MinioTestcontainerFactory INSTANCE = new MinioTestcontainerFactory();
-
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainerSingleton.java b/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainerSingleton.java
index cc2fb4f2..adceb18c 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainerSingleton.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/minio/MinioTestcontainerSingleton.java
@@ -57,4 +57,14 @@
*/
String bucketName() default "test";
+ /**
+ * Очищать содержимое bucket перед каждым тестом.
+ */
+ boolean cleanupBucket() default true;
+
+ /**
+ * Префиксы объектов, которые не нужно удалять при очистке bucket.
+ */
+ String[] excludeCleanupPrefixes() default {};
+
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchIndexCleaner.java b/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchIndexCleaner.java
new file mode 100644
index 00000000..24539f14
--- /dev/null
+++ b/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchIndexCleaner.java
@@ -0,0 +1,78 @@
+package dev.vality.testcontainers.annotations.opensearch;
+
+import com.fasterxml.jackson.core.type.TypeReference;
+import com.fasterxml.jackson.databind.ObjectMapper;
+import lombok.AccessLevel;
+import lombok.NoArgsConstructor;
+import org.apache.http.HttpHost;
+import org.apache.http.util.EntityUtils;
+import org.opensearch.client.Request;
+import org.opensearch.client.ResponseException;
+import org.opensearch.client.RestClient;
+import org.testcontainers.containers.GenericContainer;
+
+import java.io.IOException;
+import java.net.URLEncoder;
+import java.nio.charset.StandardCharsets;
+import java.util.List;
+import java.util.Map;
+
+@NoArgsConstructor(access = AccessLevel.PRIVATE)
+final class OpensearchIndexCleaner {
+
+ private static final ObjectMapper OBJECT_MAPPER = new ObjectMapper();
+ private static final TypeReference>> INDEX_LIST_TYPE = new TypeReference<>() {
+ };
+
+ static void cleanup(
+ GenericContainer> container,
+ List includedPrefixes,
+ List excludedIndexes) {
+ var builder = RestClient.builder(new HttpHost(container.getHost(), container.getFirstMappedPort(), "http"));
+ try (var client = builder.build()) {
+ for (var index : listIndexes(client)) {
+ if (shouldDelete(index, includedPrefixes, excludedIndexes)) {
+ deleteIndex(client, index);
+ }
+ }
+ } catch (IOException ex) {
+ throw new IllegalStateException("Unable to clean OpenSearch indexes", ex);
+ }
+ }
+
+ private static List listIndexes(RestClient client) throws IOException {
+ var response = client.performRequest(new Request("GET", "/_cat/indices?format=json&h=index"));
+ var entity = response.getEntity();
+ if (entity == null) {
+ return List.of();
+ }
+ try (var content = entity.getContent()) {
+ return OBJECT_MAPPER.readValue(content, INDEX_LIST_TYPE).stream()
+ .map(index -> index.get("index"))
+ .filter(name -> name != null && !name.isBlank())
+ .toList();
+ }
+ }
+
+ private static boolean shouldDelete(
+ String index,
+ List includedPrefixes,
+ List excludedIndexes) {
+ if (index.startsWith(".") || excludedIndexes.contains(index)) {
+ return false;
+ }
+ return includedPrefixes.isEmpty() || includedPrefixes.stream().anyMatch(index::startsWith);
+ }
+
+ private static void deleteIndex(RestClient client, String index) throws IOException {
+ var encodedIndex = URLEncoder.encode(index, StandardCharsets.UTF_8).replace("+", "%20");
+ try {
+ var response = client.performRequest(new Request("DELETE", "/" + encodedIndex));
+ EntityUtils.consume(response.getEntity());
+ } catch (ResponseException ex) {
+ if (ex.getResponse().getStatusLine().getStatusCode() != 404) {
+ throw ex;
+ }
+ }
+ }
+}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainer.java b/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainer.java
index e15d812a..fbed6a32 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainer.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainer.java
@@ -20,4 +20,19 @@
*/
String[] properties() default {};
+ /**
+ * Удалять тестовые индексы перед каждым тестом.
+ */
+ boolean cleanupIndexes() default true;
+
+ /**
+ * Удалять только индексы с указанными префиксами. Пустой список означает все пользовательские индексы.
+ */
+ String[] indexPrefixes() default {};
+
+ /**
+ * Имена индексов, которые не нужно удалять. Системные индексы с точкой в начале исключаются всегда.
+ */
+ String[] excludeIndexes() default {};
+
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainerExtension.java b/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainerExtension.java
index 471a470c..12c4f6e1 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainerExtension.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainerExtension.java
@@ -1,109 +1,195 @@
package dev.vality.testcontainers.annotations.opensearch;
import dev.vality.testcontainers.annotations.util.GenericContainerUtil;
-import lombok.SneakyThrows;
-import lombok.extern.slf4j.Slf4j;
-import org.apache.http.HttpHost;
+import dev.vality.testcontainers.annotations.util.SharedTestResourceLock;
+import dev.vality.testcontainers.annotations.util.TestExecutionLock;
import org.junit.jupiter.api.extension.AfterAllCallback;
import org.junit.jupiter.api.extension.BeforeAllCallback;
import org.junit.jupiter.api.extension.BeforeEachCallback;
import org.junit.jupiter.api.extension.ExtensionContext;
import org.junit.platform.commons.support.AnnotationSupport;
-import org.opensearch.client.Request;
-import org.opensearch.client.RestClient;
import org.springframework.boot.test.util.TestPropertyValues;
import org.springframework.context.ConfigurableApplicationContext;
import org.springframework.test.context.ContextConfigurationAttributes;
import org.springframework.test.context.ContextCustomizer;
import org.springframework.test.context.ContextCustomizerFactory;
+import org.springframework.test.context.MergedContextConfiguration;
import org.testcontainers.containers.GenericContainer;
+import java.util.Arrays;
import java.util.List;
import java.util.Optional;
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.ConcurrentMap;
-@Slf4j
public class OpensearchTestcontainerExtension implements BeforeAllCallback, AfterAllCallback, BeforeEachCallback {
- private static final ThreadLocal> THREAD_CONTAINER = new ThreadLocal<>();
+ private static final ConcurrentMap, ContainerReference> CONTAINERS = new ConcurrentHashMap<>();
@Override
public void beforeAll(ExtensionContext context) {
- if (findPrototypeAnnotation(context).isPresent()) {
- var container = OpensearchTestcontainerFactory.container();
- GenericContainerUtil.startContainer(container);
- THREAD_CONTAINER.set(container);
- } else if (findSingletonAnnotation(context).isPresent()) {
- var container = OpensearchTestcontainerFactory.singletonContainer();
- if (!container.isRunning()) {
- GenericContainerUtil.startContainer(container);
- }
- THREAD_CONTAINER.set(container);
- }
+ var testClass = context.getRequiredTestClass();
+ getOrStart(testClass, configuration(testClass));
}
@Override
- @SneakyThrows
public void beforeEach(ExtensionContext context) {
- var container = THREAD_CONTAINER.get();
- if (container != null && container.isRunning()) {
- var builder = RestClient.builder(new HttpHost(container.getHost(), container.getFirstMappedPort()));
- try (var client = builder.build()) {
- var deleteRequest = new Request("DELETE", "/*");
- client.performRequest(deleteRequest);
+ var testClass = context.getRequiredTestClass();
+ var configuration = configuration(testClass);
+ if (configuration.cleanupIndexes()) {
+ TestExecutionLock.acquire(context);
+ try {
+ var reference = getOrStart(testClass, configuration);
+ OpensearchIndexCleaner.cleanup(
+ reference.container(),
+ configuration.indexPrefixes(),
+ configuration.excludedIndexes());
+ } catch (RuntimeException ex) {
+ TestExecutionLock.release(context);
+ throw ex;
}
}
}
@Override
public void afterAll(ExtensionContext context) {
- if (findPrototypeAnnotation(context).isPresent()) {
- var container = THREAD_CONTAINER.get();
- if (container != null && container.isRunning()) {
- container.stop();
+ var testClass = context.getRequiredTestClass();
+ var reference = CONTAINERS.remove(testClass);
+ try {
+ if (reference != null && !reference.singleton()) {
+ reference.container().stop();
+ }
+ } finally {
+ if (reference != null && reference.singleton()) {
+ SharedTestResourceLock.release(testClass);
}
- THREAD_CONTAINER.remove();
- } else if (findSingletonAnnotation(context).isPresent()) {
- THREAD_CONTAINER.remove();
}
}
- private static Optional findPrototypeAnnotation(ExtensionContext context) {
- return AnnotationSupport.findAnnotation(context.getTestClass(), OpensearchTestcontainer.class);
+ private static ContainerReference getOrStart(Class> testClass, Configuration configuration) {
+ return CONTAINERS.computeIfAbsent(testClass, ignored -> createAndStart(testClass, configuration));
}
- private static Optional findPrototypeAnnotation(Class> testClass) {
- return AnnotationSupport.findAnnotation(testClass, OpensearchTestcontainer.class);
+ private static ContainerReference createAndStart(Class> testClass, Configuration configuration) {
+ if (configuration.singleton()) {
+ SharedTestResourceLock.acquire(testClass);
+ }
+ try {
+ var container = configuration.singleton()
+ ? OpensearchTestcontainerFactory.singletonContainer()
+ : OpensearchTestcontainerFactory.container();
+ try {
+ GenericContainerUtil.startContainer(container);
+ if (configuration.singleton() && configuration.cleanupIndexes()) {
+ OpensearchIndexCleaner.cleanup(
+ container,
+ configuration.indexPrefixes(),
+ configuration.excludedIndexes());
+ }
+ return new ContainerReference(container, configuration.singleton());
+ } catch (RuntimeException ex) {
+ try {
+ if (configuration.singleton()) {
+ OpensearchTestcontainerFactory.discardSingleton(container);
+ } else {
+ container.stop();
+ }
+ } catch (RuntimeException stopException) {
+ ex.addSuppressed(stopException);
+ }
+ throw ex;
+ }
+ } catch (RuntimeException ex) {
+ if (configuration.singleton()) {
+ SharedTestResourceLock.release(testClass);
+ }
+ throw ex;
+ }
}
- private static Optional findSingletonAnnotation(ExtensionContext context) {
- return AnnotationSupport.findAnnotation(context.getTestClass(), OpensearchTestcontainerSingleton.class);
+ private static Configuration configuration(Class> testClass) {
+ var prototype = findPrototypeAnnotation(testClass);
+ if (prototype.isPresent()) {
+ var annotation = prototype.get();
+ return new Configuration(
+ false,
+ annotation.cleanupIndexes(),
+ List.copyOf(Arrays.asList(annotation.indexPrefixes())),
+ List.copyOf(Arrays.asList(annotation.excludeIndexes())),
+ List.copyOf(Arrays.asList(annotation.properties())));
+ }
+ var annotation = findSingletonAnnotation(testClass)
+ .orElseThrow(() -> new IllegalStateException("OpenSearch test annotation not found"));
+ return new Configuration(
+ true,
+ annotation.cleanupIndexes(),
+ List.copyOf(Arrays.asList(annotation.indexPrefixes())),
+ List.copyOf(Arrays.asList(annotation.excludeIndexes())),
+ List.copyOf(Arrays.asList(annotation.properties())));
+ }
+
+ private static Optional findPrototypeAnnotation(Class> testClass) {
+ return AnnotationSupport.findAnnotation(testClass, OpensearchTestcontainer.class);
}
private static Optional findSingletonAnnotation(Class> testClass) {
return AnnotationSupport.findAnnotation(testClass, OpensearchTestcontainerSingleton.class);
}
+ private static final class ContainerReference {
+
+ private final GenericContainer> container;
+ private final boolean singleton;
+
+ private ContainerReference(GenericContainer> container, boolean singleton) {
+ this.container = container;
+ this.singleton = singleton;
+ }
+
+ private GenericContainer> container() {
+ return container;
+ }
+
+ private boolean singleton() {
+ return singleton;
+ }
+
+ }
+
+ private record Configuration(
+ boolean singleton,
+ boolean cleanupIndexes,
+ List indexPrefixes,
+ List excludedIndexes,
+ List properties) {
+ }
+
public static class OpensearchTestcontainerContextCustomizerFactory implements ContextCustomizerFactory {
@Override
public ContextCustomizer createContextCustomizer(
Class> testClass,
List configAttributes) {
- return (context, mergedConfig) -> {
- if (findPrototypeAnnotation(testClass).isPresent()) {
- init(context, findPrototypeAnnotation(testClass).get().properties());
- } else if (findSingletonAnnotation(testClass).isPresent()) {
- init(context, findSingletonAnnotation(testClass).get().properties());
- }
- };
+ if (findPrototypeAnnotation(testClass).isEmpty() && findSingletonAnnotation(testClass).isEmpty()) {
+ return null;
+ }
+ return new OpensearchContextCustomizer(testClass, configuration(testClass));
}
+ }
- private void init(ConfigurableApplicationContext context, String[] properties) {
- var container = THREAD_CONTAINER.get();
+ private record OpensearchContextCustomizer(
+ Class> testClass,
+ Configuration configuration) implements ContextCustomizer {
+
+ @Override
+ public void customizeContext(
+ ConfigurableApplicationContext context,
+ MergedContextConfiguration mergedConfig) {
+ var container = getOrStart(testClass, configuration).container();
TestPropertyValues.of(
"opensearch.hostname=" + container.getHost(),
"opensearch.port=" + container.getFirstMappedPort())
- .and(properties)
+ .and(configuration.properties().toArray(String[]::new))
.applyTo(context);
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainerFactory.java b/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainerFactory.java
index 37a9859b..ef6be9f8 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainerFactory.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainerFactory.java
@@ -1,15 +1,12 @@
package dev.vality.testcontainers.annotations.opensearch;
+import dev.vality.testcontainers.annotations.util.ContainerShutdownRegistry;
import lombok.AccessLevel;
import lombok.NoArgsConstructor;
-import lombok.Synchronized;
import org.testcontainers.containers.GenericContainer;
-import org.testcontainers.containers.Network;
import org.testcontainers.containers.wait.strategy.HttpWaitStrategy;
import org.testcontainers.utility.DockerImageName;
-import java.util.UUID;
-
import static dev.vality.testcontainers.annotations.util.SpringApplicationPropertiesLoader.loadDefaultLibraryProperty;
@NoArgsConstructor(access = AccessLevel.PRIVATE)
@@ -18,7 +15,7 @@ public class OpensearchTestcontainerFactory {
private static final String OPENSEARCH_IMAGE_NAME = "opensearchproject/opensearch";
private static final String TAG_PROPERTY = "testcontainers.opensearch.tag";
- private GenericContainer> opensearchContainer;
+ private volatile GenericContainer> singletonContainer;
public static GenericContainer> container() {
return instance().create();
@@ -28,39 +25,53 @@ public static GenericContainer> singletonContainer() {
return instance().getOrCreateSingletonContainer();
}
+ static void discardSingleton(GenericContainer> container) {
+ instance().discard(container);
+ }
+
private static OpensearchTestcontainerFactory instance() {
return SingletonHolder.INSTANCE;
}
- @Synchronized
+ private void discard(GenericContainer> container) {
+ synchronized (this) {
+ if (singletonContainer != container) {
+ return;
+ }
+ singletonContainer = null;
+ }
+ ContainerShutdownRegistry.unregister(container);
+ container.stop();
+ }
+
private GenericContainer> getOrCreateSingletonContainer() {
- if (opensearchContainer != null) {
- return opensearchContainer;
+ var result = singletonContainer;
+ if (result == null) {
+ synchronized (this) {
+ result = singletonContainer;
+ if (result == null) {
+ result = ContainerShutdownRegistry.register(create());
+ singletonContainer = result;
+ }
+ }
}
- opensearchContainer = create();
- return opensearchContainer;
+ return result;
}
private GenericContainer> create() {
- var container = new GenericContainer<>(
- DockerImageName
- .parse(OPENSEARCH_IMAGE_NAME)
- .withTag(loadDefaultLibraryProperty(TAG_PROPERTY)));
- container.withNetworkAliases("opensearch-" + UUID.randomUUID());
- container.withNetwork(Network.SHARED);
- container.withExposedPorts(9200, 9600);
- container.setWaitStrategy((new HttpWaitStrategy())
- .forPort(9200)
- .forStatusCodeMatching(response -> response == 200 || response == 401));
- container.withEnv("discovery.type", "single-node");
- container.withEnv("DISABLE_INSTALL_DEMO_CONFIG", "true");
- container.withEnv("DISABLE_SECURITY_PLUGIN", "true");
- return container;
+ return new GenericContainer<>(DockerImageName.parse(OPENSEARCH_IMAGE_NAME)
+ .withTag(loadDefaultLibraryProperty(TAG_PROPERTY)))
+ .withExposedPorts(9200, 9600)
+ .waitingFor(new HttpWaitStrategy()
+ .forPort(9200)
+ .forStatusCodeMatching(response -> response == 200 || response == 401))
+ .withEnv("discovery.type", "single-node")
+ .withEnv("DISABLE_INSTALL_DEMO_CONFIG", "true")
+ .withEnv("DISABLE_SECURITY_PLUGIN", "true");
}
private static class SingletonHolder {
private static final OpensearchTestcontainerFactory INSTANCE = new OpensearchTestcontainerFactory();
-
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainerSingleton.java b/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainerSingleton.java
index b1ad863b..6d9c6720 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainerSingleton.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/opensearch/OpensearchTestcontainerSingleton.java
@@ -22,4 +22,19 @@
*/
String[] properties() default {};
+ /**
+ * Удалять тестовые индексы перед каждым тестом.
+ */
+ boolean cleanupIndexes() default true;
+
+ /**
+ * Удалять только индексы с указанными префиксами. Пустой список означает все пользовательские индексы.
+ */
+ String[] indexPrefixes() default {};
+
+ /**
+ * Имена индексов, которые не нужно удалять. Системные индексы с точкой в начале исключаются всегда.
+ */
+ String[] excludeIndexes() default {};
+
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/postgresql/EmbeddedPostgresqlTestContextCustomizerFactory.java b/src/main/java/dev/vality/testcontainers/annotations/postgresql/EmbeddedPostgresqlTestContextCustomizerFactory.java
index 1b766aff..7ba7714e 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/postgresql/EmbeddedPostgresqlTestContextCustomizerFactory.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/postgresql/EmbeddedPostgresqlTestContextCustomizerFactory.java
@@ -5,7 +5,9 @@
import org.springframework.test.context.ContextConfigurationAttributes;
import org.springframework.test.context.ContextCustomizer;
import org.springframework.test.context.ContextCustomizerFactory;
+import org.springframework.test.context.MergedContextConfiguration;
+import java.util.Arrays;
import java.util.List;
public class EmbeddedPostgresqlTestContextCustomizerFactory implements ContextCustomizerFactory {
@@ -14,32 +16,48 @@ public class EmbeddedPostgresqlTestContextCustomizerFactory implements ContextCu
public ContextCustomizer createContextCustomizer(
Class> testClass,
List configAttributes) {
- return (context, mergedConfig) ->
- EmbeddedPostgresqlTestExtension.findAnnotation(testClass)
- .ifPresent(annotation -> init(context, annotation));
+ return EmbeddedPostgresqlTestExtension.findAnnotation(testClass)
+ .map(annotation -> new EmbeddedPostgresqlContextCustomizer(
+ testClass,
+ annotation.database(),
+ annotation.username(),
+ annotation.password(),
+ List.copyOf(Arrays.asList(annotation.properties()))))
+ .orElse(null);
}
- private void init(ConfigurableApplicationContext context, EmbeddedPostgresqlTest annotation) {
- var postgresql = EmbeddedPostgresqlTestExtension.getOrStart(annotation);
- var jdbcUrl = postgresql.jdbcUrl();
- var username = annotation.username();
- var password = annotation.password();
- TestPropertyValues.of(
- "spring.datasource.url=" + jdbcUrl,
- "spring.datasource.username=" + username,
- "spring.datasource.password=" + password,
- "spring.flyway.url=" + jdbcUrl,
- "spring.flyway.user=" + username,
- "spring.flyway.password=" + password,
- "postgres.db.url=" + jdbcUrl,
- "postgres.db.user=" + username,
- "postgres.db.username=" + username,
- "postgres.db.password=" + password,
- "flyway.url=" + jdbcUrl,
- "flyway.user=" + username,
- "flyway.password=" + password,
- "flyway.postgresql.transactional.lock=false")
- .and(annotation.properties())
- .applyTo(context);
+ private record EmbeddedPostgresqlContextCustomizer(
+ Class> testClass,
+ String database,
+ String username,
+ String password,
+ List properties) implements ContextCustomizer {
+
+ @Override
+ public void customizeContext(
+ ConfigurableApplicationContext context,
+ MergedContextConfiguration mergedConfig) {
+ var annotation = EmbeddedPostgresqlTestExtension.findAnnotation(testClass)
+ .orElseThrow(() -> new IllegalStateException("Embedded PostgreSQL annotation not found"));
+ var postgresql = EmbeddedPostgresqlTestExtension.getOrStart(testClass, annotation);
+ var jdbcUrl = postgresql.jdbcUrl();
+ TestPropertyValues.of(
+ "spring.datasource.url=" + jdbcUrl,
+ "spring.datasource.username=" + username,
+ "spring.datasource.password=" + password,
+ "spring.flyway.url=" + jdbcUrl,
+ "spring.flyway.user=" + username,
+ "spring.flyway.password=" + password,
+ "postgres.db.url=" + jdbcUrl,
+ "postgres.db.user=" + username,
+ "postgres.db.username=" + username,
+ "postgres.db.password=" + password,
+ "flyway.url=" + jdbcUrl,
+ "flyway.user=" + username,
+ "flyway.password=" + password,
+ "flyway.postgresql.transactional.lock=false")
+ .and(properties.toArray(String[]::new))
+ .applyTo(context);
+ }
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/postgresql/EmbeddedPostgresqlTestExtension.java b/src/main/java/dev/vality/testcontainers/annotations/postgresql/EmbeddedPostgresqlTestExtension.java
index d3fe7392..8a3e4998 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/postgresql/EmbeddedPostgresqlTestExtension.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/postgresql/EmbeddedPostgresqlTestExtension.java
@@ -1,80 +1,168 @@
package dev.vality.testcontainers.annotations.postgresql;
+import dev.vality.testcontainers.annotations.util.TestExecutionLock;
import io.zonky.test.db.postgres.embedded.EmbeddedPostgres;
-import lombok.SneakyThrows;
import org.junit.jupiter.api.extension.AfterAllCallback;
import org.junit.jupiter.api.extension.BeforeAllCallback;
import org.junit.jupiter.api.extension.BeforeEachCallback;
import org.junit.jupiter.api.extension.ExtensionContext;
import org.junit.platform.commons.support.AnnotationSupport;
+import java.sql.DriverManager;
+import java.sql.SQLException;
+import java.util.Arrays;
import java.util.List;
import java.util.Optional;
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.ConcurrentMap;
public class EmbeddedPostgresqlTestExtension implements BeforeAllCallback, BeforeEachCallback, AfterAllCallback {
- private static final ThreadLocal THREAD_POSTGRESQL = new ThreadLocal<>();
+ private static final ConcurrentMap, EmbeddedPostgresql> INSTANCES = new ConcurrentHashMap<>();
@Override
public void beforeAll(ExtensionContext context) {
- findAnnotation(context).ifPresent(annotation -> THREAD_POSTGRESQL.set(getOrStart(annotation)));
+ var testClass = context.getRequiredTestClass();
+ findAnnotation(testClass).ifPresent(annotation -> getOrStart(testClass, annotation));
}
@Override
public void beforeEach(ExtensionContext context) {
- findAnnotation(context).ifPresent(annotation -> {
+ var testClass = context.getRequiredTestClass();
+ findAnnotation(testClass).ifPresent(annotation -> {
if (annotation.truncateTables()) {
- var postgresql = getOrStart(annotation);
- PostgresqlDatabaseCleaner.cleanupDatabaseTables(
- postgresql.jdbcUrl(),
- annotation.username(),
- annotation.password(),
- List.of(annotation.excludeTruncateTables()));
+ TestExecutionLock.acquire(context);
+ try {
+ var postgresql = getOrStart(testClass, annotation);
+ PostgresqlDatabaseCleaner.cleanupDatabaseTables(
+ postgresql.jdbcUrl(),
+ annotation.username(),
+ annotation.password(),
+ List.copyOf(Arrays.asList(annotation.excludeTruncateTables())));
+ } catch (RuntimeException ex) {
+ TestExecutionLock.release(context);
+ throw ex;
+ }
}
});
}
@Override
public void afterAll(ExtensionContext context) {
- findAnnotation(context).ifPresent(annotation -> {
- var postgresql = THREAD_POSTGRESQL.get();
- THREAD_POSTGRESQL.remove();
- if (postgresql != null) {
- postgresql.close();
- }
- });
- }
-
- private static Optional findAnnotation(ExtensionContext context) {
- return AnnotationSupport.findAnnotation(context.getTestClass(), EmbeddedPostgresqlTest.class);
+ var postgresql = INSTANCES.remove(context.getRequiredTestClass());
+ if (postgresql != null) {
+ postgresql.close();
+ }
}
static Optional findAnnotation(Class> testClass) {
return AnnotationSupport.findAnnotation(testClass, EmbeddedPostgresqlTest.class);
}
- static EmbeddedPostgresql getOrStart(EmbeddedPostgresqlTest annotation) {
- var postgresql = THREAD_POSTGRESQL.get();
- if (postgresql == null) {
- postgresql = EmbeddedPostgresql.start(annotation);
- THREAD_POSTGRESQL.set(postgresql);
- }
- return postgresql;
+ static EmbeddedPostgresql getOrStart(Class> testClass, EmbeddedPostgresqlTest annotation) {
+ return INSTANCES.computeIfAbsent(testClass, ignored -> EmbeddedPostgresql.start(annotation));
}
record EmbeddedPostgresql(EmbeddedPostgres delegate, String jdbcUrl) {
- @SneakyThrows
private static EmbeddedPostgresql start(EmbeddedPostgresqlTest annotation) {
- var postgres = EmbeddedPostgres.start();
- return new EmbeddedPostgresql(
- postgres,
- postgres.getJdbcUrl(annotation.database(), annotation.username()));
+ EmbeddedPostgres postgres = null;
+ try {
+ postgres = EmbeddedPostgres.start();
+ initializeDatabase(postgres, annotation);
+ return new EmbeddedPostgresql(
+ postgres,
+ postgres.getJdbcUrl(annotation.username(), annotation.database()));
+ } catch (Exception ex) {
+ if (postgres != null) {
+ try {
+ postgres.close();
+ } catch (Exception closeException) {
+ ex.addSuppressed(closeException);
+ }
+ }
+ throw new IllegalStateException("Unable to start embedded PostgreSQL", ex);
+ }
+ }
+
+ private static void initializeDatabase(
+ EmbeddedPostgres postgres,
+ EmbeddedPostgresqlTest annotation) throws SQLException {
+ var adminUrl = postgres.getJdbcUrl("postgres", "postgres");
+ try (var connection = DriverManager.getConnection(adminUrl, "postgres", "")) {
+ ensureRole(connection, annotation.username(), annotation.password());
+ ensureDatabase(connection, annotation.database(), annotation.username());
+ }
+ }
+
+ private static void ensureRole(
+ java.sql.Connection connection,
+ String username,
+ String password) throws SQLException {
+ var roleExists = false;
+ try (var query = connection.prepareStatement("SELECT 1 FROM pg_roles WHERE rolname = ?")) {
+ query.setString(1, username);
+ try (var result = query.executeQuery()) {
+ roleExists = result.next();
+ }
+ }
+
+ try (var statement = connection.createStatement()) {
+ if (!roleExists) {
+ var passwordClause = password.isEmpty()
+ ? ""
+ : " PASSWORD '" + quoteLiteral(password) + "'";
+ statement.execute("CREATE ROLE " + quoteIdentifier(username) + " LOGIN" + passwordClause);
+ } else if (!password.isEmpty()) {
+ statement.execute("ALTER ROLE " + quoteIdentifier(username) +
+ " PASSWORD '" + quoteLiteral(password) + "'");
+ }
+ }
+ }
+
+ private static void ensureDatabase(
+ java.sql.Connection connection,
+ String database,
+ String owner) throws SQLException {
+ String currentOwner = null;
+ try (var query = connection.prepareStatement(
+ "SELECT pg_get_userbyid(datdba) AS owner_name FROM pg_database WHERE datname = ?")) {
+ query.setString(1, database);
+ try (var result = query.executeQuery()) {
+ if (result.next()) {
+ currentOwner = result.getString("owner_name");
+ }
+ }
+ }
+
+ try (var statement = connection.createStatement()) {
+ if (currentOwner == null) {
+ statement.execute("CREATE DATABASE " + quoteIdentifier(database) +
+ " OWNER " + quoteIdentifier(owner));
+ } else if (!currentOwner.equals(owner)) {
+ statement.execute("ALTER DATABASE " + quoteIdentifier(database) +
+ " OWNER TO " + quoteIdentifier(owner));
+ }
+ }
+ }
+
+ private static String quoteIdentifier(String identifier) {
+ if (identifier == null || identifier.isBlank()) {
+ throw new IllegalArgumentException("PostgreSQL identifier must not be blank");
+ }
+ return '"' + identifier.replace("\"", "\"\"") + '"';
+ }
+
+ private static String quoteLiteral(String value) {
+ return value.replace("'", "''");
}
- @SneakyThrows
private void close() {
- delegate.close();
+ try {
+ delegate.close();
+ } catch (Exception ex) {
+ throw new IllegalStateException("Unable to stop embedded PostgreSQL", ex);
+ }
}
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlContainerExtension.java b/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlContainerExtension.java
index 2ab2d5a9..1c70bd67 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlContainerExtension.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlContainerExtension.java
@@ -1,32 +1,27 @@
package dev.vality.testcontainers.annotations.postgresql;
-import lombok.SneakyThrows;
-import lombok.extern.slf4j.Slf4j;
-import org.testcontainers.containers.Network;
import org.testcontainers.containers.PostgreSQLContainer;
import org.testcontainers.utility.DockerImageName;
import java.util.List;
-import java.util.UUID;
import static dev.vality.testcontainers.annotations.util.SpringApplicationPropertiesLoader.loadDefaultLibraryProperty;
-@Slf4j
public class PostgresqlContainerExtension extends PostgreSQLContainer {
private static final String POSTGRESQL_IMAGE_NAME = "postgres";
private static final String TAG_PROPERTY = "testcontainers.postgresql.tag";
public PostgresqlContainerExtension() {
- super(DockerImageName
- .parse(POSTGRESQL_IMAGE_NAME)
+ super(DockerImageName.parse(POSTGRESQL_IMAGE_NAME)
.withTag(loadDefaultLibraryProperty(TAG_PROPERTY)));
- withNetworkAliases("postgresql-" + UUID.randomUUID());
- withNetwork(Network.SHARED);
}
- @SneakyThrows
public void cleanupDatabaseTables(List excludedTables) {
- PostgresqlDatabaseCleaner.cleanupDatabaseTables(getJdbcUrl(), getUsername(), getPassword(), excludedTables);
+ PostgresqlDatabaseCleaner.cleanupDatabaseTables(
+ getJdbcUrl(),
+ getUsername(),
+ getPassword(),
+ excludedTables);
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlDatabaseCleaner.java b/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlDatabaseCleaner.java
index 530db85d..a1fa09ef 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlDatabaseCleaner.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlDatabaseCleaner.java
@@ -2,54 +2,81 @@
import lombok.AccessLevel;
import lombok.NoArgsConstructor;
-import lombok.SneakyThrows;
import lombok.extern.slf4j.Slf4j;
import java.sql.Connection;
import java.sql.DriverManager;
-import java.util.ArrayList;
-import java.util.HashSet;
-import java.util.List;
-import java.util.Set;
+import java.sql.SQLException;
+import java.util.*;
+import java.util.stream.Collectors;
@Slf4j
@NoArgsConstructor(access = AccessLevel.PRIVATE)
public class PostgresqlDatabaseCleaner {
- private static final String CURRENT_SCHEMA_QUERY = "SELECT schema_name FROM information_schema.schemata";
- private static final String TABLES_QUERY = "SELECT tablename FROM pg_tables " +
- "WHERE schemaname = ? AND tablename NOT LIKE 'flyway%'AND tablename NOT LIKE 'schema_version'";
- private static final String TRUNCATE_TABLE_QUERY = "TRUNCATE TABLE %s.%s CASCADE";
- private static final Set EXCLUDE_SCHEMAS = Set.of("information_schema");
- private static final String PG_ = "pg_";
- private static final String SQL_ = "sql_";
+ private static final String SCHEMAS_QUERY = "SELECT schema_name FROM information_schema.schemata";
+ private static final String TABLES_QUERY = """
+ SELECT tablename
+ FROM pg_tables
+ WHERE schemaname = ?
+ AND tablename NOT LIKE 'flyway%'
+ AND tablename <> 'schema_version'
+ """;
+ private static final Set EXCLUDED_SCHEMAS = Set.of("information_schema");
- @SneakyThrows
public static void cleanupDatabaseTables(
String jdbcUrl,
String username,
String password,
List excludedTables) {
try (var connection = DriverManager.getConnection(jdbcUrl, username, password)) {
+ cleanup(connection, excludedTables == null ? List.of() : excludedTables);
+ } catch (SQLException ex) {
+ throw new IllegalStateException("Unable to clean PostgreSQL database " + jdbcUrl, ex);
+ }
+ }
+
+ private static void cleanup(Connection connection, List excludedTables) throws SQLException {
+ var previousAutoCommit = connection.getAutoCommit();
+ Throwable failure = null;
+ connection.setAutoCommit(false);
+ try {
+ var exclusions = new HashSet<>(excludedTables);
+ var tables = new ArrayList();
for (var schema : getSchemas(connection)) {
- var tables = getUserTables(connection, schema, excludedTables);
- if (!tables.isEmpty()) {
- truncateTables(connection, schema, tables);
+ tables.addAll(getUserTables(connection, schema, exclusions));
+ }
+ truncateTables(connection, tables);
+ connection.commit();
+ } catch (SQLException | RuntimeException ex) {
+ failure = ex;
+ try {
+ connection.rollback();
+ } catch (SQLException rollbackException) {
+ ex.addSuppressed(rollbackException);
+ }
+ throw ex;
+ } finally {
+ try {
+ connection.setAutoCommit(previousAutoCommit);
+ } catch (SQLException restoreException) {
+ if (failure != null) {
+ failure.addSuppressed(restoreException);
+ } else {
+ throw restoreException;
}
}
}
}
- @SneakyThrows
- private static Set getSchemas(Connection connection) {
- var schemas = new HashSet();
- try (
- var statement = connection.createStatement();
- var resultSet = statement.executeQuery(CURRENT_SCHEMA_QUERY)) {
+ private static Set getSchemas(Connection connection) throws SQLException {
+ var schemas = new LinkedHashSet();
+ try (var statement = connection.createStatement(); var resultSet = statement.executeQuery(SCHEMAS_QUERY)) {
while (resultSet.next()) {
var schema = resultSet.getString("schema_name");
- if (!EXCLUDE_SCHEMAS.contains(schema)
- && !schema.startsWith(PG_) && !schema.startsWith(SQL_)) {
+ if (!EXCLUDED_SCHEMAS.contains(schema)
+ && !schema.startsWith("pg_")
+ && !schema.startsWith("sql_")) {
schemas.add(schema);
}
}
@@ -57,19 +84,18 @@ private static Set getSchemas(Connection connection) {
return schemas;
}
- @SneakyThrows
- private static List getUserTables(Connection connection, String schema, List excludedTables) {
- var tables = new ArrayList();
+ private static List getUserTables(
+ Connection connection,
+ String schema,
+ Set exclusions) throws SQLException {
+ var tables = new ArrayList();
try (var statement = connection.prepareStatement(TABLES_QUERY)) {
statement.setString(1, schema);
try (var resultSet = statement.executeQuery()) {
while (resultSet.next()) {
- var tableName = resultSet.getString("tablename");
- boolean isTruncatable = !tableName.startsWith(PG_)
- && !tableName.startsWith(SQL_)
- && !excludedTables.contains(tableName);
- if (isTruncatable) {
- tables.add(tableName);
+ var table = resultSet.getString("tablename");
+ if (!isExcluded(schema, table, exclusions)) {
+ tables.add(new TableReference(schema, table));
}
}
}
@@ -77,15 +103,28 @@ private static List getUserTables(Connection connection, String schema,
return tables;
}
- @SneakyThrows
- private static void truncateTables(Connection connection, String schema, List tables) {
+ private static boolean isExcluded(String schema, String table, Set exclusions) {
+ return exclusions.contains(table) || exclusions.contains(schema + "." + table);
+ }
+
+ private static void truncateTables(Connection connection, List tables) throws SQLException {
+ if (tables.isEmpty()) {
+ return;
+ }
+ var qualifiedTables = tables.stream()
+ .map(table -> quoteIdentifier(table.schema()) + "." + quoteIdentifier(table.table()))
+ .collect(Collectors.joining(", "));
+ var sql = "TRUNCATE TABLE " + qualifiedTables + " RESTART IDENTITY CASCADE";
+ log.debug("Cleaning PostgreSQL tables: {}", qualifiedTables);
try (var statement = connection.createStatement()) {
- statement.execute("SET session_replication_role = 'replica'");
- for (var table : tables) {
- log.debug("Truncating table: {}", table);
- statement.execute(String.format(TRUNCATE_TABLE_QUERY, schema, table));
- }
- statement.execute("SET session_replication_role = 'origin'");
+ statement.execute(sql);
}
}
+
+ private static String quoteIdentifier(String identifier) {
+ return '"' + identifier.replace("\"", "\"\"") + '"';
+ }
+
+ private record TableReference(String schema, String table) {
+ }
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlTestcontainerExtension.java b/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlTestcontainerExtension.java
index e52e5a31..32cf230d 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlTestcontainerExtension.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlTestcontainerExtension.java
@@ -1,7 +1,8 @@
package dev.vality.testcontainers.annotations.postgresql;
import dev.vality.testcontainers.annotations.util.GenericContainerUtil;
-import lombok.extern.slf4j.Slf4j;
+import dev.vality.testcontainers.annotations.util.SharedTestResourceLock;
+import dev.vality.testcontainers.annotations.util.TestExecutionLock;
import org.junit.jupiter.api.extension.AfterAllCallback;
import org.junit.jupiter.api.extension.BeforeAllCallback;
import org.junit.jupiter.api.extension.BeforeEachCallback;
@@ -12,123 +13,183 @@
import org.springframework.test.context.ContextConfigurationAttributes;
import org.springframework.test.context.ContextCustomizer;
import org.springframework.test.context.ContextCustomizerFactory;
+import org.springframework.test.context.MergedContextConfiguration;
+import java.util.Arrays;
import java.util.List;
import java.util.Optional;
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.ConcurrentMap;
-/**
- * {@code @PostgresqlTestcontainerExtension} инициализирует тестконтейнер из {@link PostgresqlTestcontainerFactory},
- * настраивает, стартует, валидирует и останавливает
- *
{@link PostgresqlTestcontainerContextCustomizerFactory}
- * Инициализация настроек контейнеров в спринговый контекст тестового приложения реализован
- * под капотом аннотаций, на уровне реализации интерфейса —
- * информация о настройках используемого тестконтейнера и передаваемые через параметры аннотации настройки
- * инициализируются через {@link TestPropertyValues} и сливаются с текущим получаемым контекстом
- * приложения {@link ConfigurableApplicationContext}
- *
Инициализация кастомизированных фабрик с инициализацией настроек осуществляется через описание бинов
- * в файле META-INF/spring.factories
- *
- * @see PostgresqlTestcontainerFactory PostgresqlTestcontainerFactory
- * @see PostgresqlTestcontainerContextCustomizerFactory PostgresqlTestcontainerContextCustomizerFactory
- * @see TestPropertyValues TestPropertyValues
- * @see ConfigurableApplicationContext ConfigurableApplicationContext
- * @see BeforeAllCallback BeforeAllCallback
- * @see AfterAllCallback AfterAllCallback
- */
-@Slf4j
public class PostgresqlTestcontainerExtension implements BeforeAllCallback, AfterAllCallback, BeforeEachCallback {
- private static final ThreadLocal THREAD_CONTAINER = new ThreadLocal<>();
+ private static final ConcurrentMap, ContainerReference> CONTAINERS = new ConcurrentHashMap<>();
@Override
public void beforeAll(ExtensionContext context) {
- if (findPrototypeAnnotation(context).isPresent()) {
- var container = PostgresqlTestcontainerFactory.container();
- GenericContainerUtil.startContainer(container);
- THREAD_CONTAINER.set(container);
- } else if (findSingletonAnnotation(context).isPresent()) {
- var container = PostgresqlTestcontainerFactory.singletonContainer();
- if (!container.isRunning()) {
- GenericContainerUtil.startContainer(container);
- } else if (findSingletonAnnotation(context).get().truncateTables()) {
- var excludedTables = Optional.ofNullable(findSingletonAnnotation(context).get().excludeTruncateTables())
- .map(List::of)
- .orElse(List.of());
- container.cleanupDatabaseTables(excludedTables);
- }
-
- THREAD_CONTAINER.set(container);
- }
+ getOrStart(context.getRequiredTestClass());
}
@Override
public void beforeEach(ExtensionContext context) {
- var container = THREAD_CONTAINER.get();
- if (findPrototypeAnnotation(context).isPresent()) {
- var annotation = findPrototypeAnnotation(context).get();
- if (container != null && container.isRunning() && annotation.truncateTables()) {
- var excludedTables = Optional.ofNullable(annotation.excludeTruncateTables())
- .map(List::of)
- .orElse(List.of());
- container.cleanupDatabaseTables(excludedTables);
- }
- } else if (findSingletonAnnotation(context).isPresent()) {
- var annotation = findSingletonAnnotation(context).get();
- if (container != null && container.isRunning() && annotation.truncateTables()) {
- var excludedTables = Optional.ofNullable(annotation.excludeTruncateTables())
- .map(List::of)
- .orElse(List.of());
- container.cleanupDatabaseTables(excludedTables);
+ var testClass = context.getRequiredTestClass();
+ findCleanupConfiguration(testClass).ifPresent(configuration -> {
+ if (configuration.truncateTables()) {
+ TestExecutionLock.acquire(context);
+ try {
+ var reference = getOrStart(testClass);
+ reference.container().cleanupDatabaseTables(configuration.excludedTables());
+ } catch (RuntimeException ex) {
+ TestExecutionLock.release(context);
+ throw ex;
+ }
}
- }
+ });
}
@Override
public void afterAll(ExtensionContext context) {
- if (findPrototypeAnnotation(context).isPresent()) {
- var container = THREAD_CONTAINER.get();
- if (container != null && container.isRunning()) {
- container.stop();
+ var testClass = context.getRequiredTestClass();
+ var reference = CONTAINERS.remove(testClass);
+ try {
+ if (reference != null && !reference.singleton()) {
+ reference.container().stop();
+ }
+ } finally {
+ if (reference != null && reference.singleton()) {
+ SharedTestResourceLock.release(testClass);
}
- THREAD_CONTAINER.remove();
- } else if (findSingletonAnnotation(context).isPresent()) {
- THREAD_CONTAINER.remove();
}
}
- private static Optional findPrototypeAnnotation(ExtensionContext context) {
- return AnnotationSupport.findAnnotation(context.getTestClass(), PostgresqlTestcontainer.class);
+ private static ContainerReference getOrStart(Class> testClass) {
+ return CONTAINERS.computeIfAbsent(testClass, PostgresqlTestcontainerExtension::createAndStart);
}
- private static Optional findPrototypeAnnotation(Class> testClass) {
- return AnnotationSupport.findAnnotation(testClass, PostgresqlTestcontainer.class);
+ private static ContainerReference createAndStart(Class> testClass) {
+ var prototype = findPrototypeAnnotation(testClass);
+ var singletonAnnotation = findSingletonAnnotation(testClass);
+ if (prototype.isEmpty() && singletonAnnotation.isEmpty()) {
+ throw new IllegalStateException("PostgreSQL test annotation not found");
+ }
+
+ var singleton = singletonAnnotation.isPresent();
+ if (singleton) {
+ SharedTestResourceLock.acquire(testClass);
+ }
+ try {
+ var container = singleton
+ ? PostgresqlTestcontainerFactory.singletonContainer()
+ : PostgresqlTestcontainerFactory.container();
+ try {
+ GenericContainerUtil.startContainer(container);
+ if (singleton) {
+ findCleanupConfiguration(testClass).ifPresent(configuration -> {
+ if (configuration.truncateTables()) {
+ container.cleanupDatabaseTables(configuration.excludedTables());
+ }
+ });
+ }
+ return new ContainerReference(container, singleton);
+ } catch (RuntimeException ex) {
+ try {
+ if (singleton) {
+ PostgresqlTestcontainerFactory.discardSingleton(container);
+ } else {
+ container.stop();
+ }
+ } catch (RuntimeException stopException) {
+ ex.addSuppressed(stopException);
+ }
+ throw ex;
+ }
+ } catch (RuntimeException ex) {
+ if (singleton) {
+ SharedTestResourceLock.release(testClass);
+ }
+ throw ex;
+ }
+ }
+
+ private static Optional findCleanupConfiguration(Class> testClass) {
+ var prototype = findPrototypeAnnotation(testClass);
+ if (prototype.isPresent()) {
+ var annotation = prototype.get();
+ return Optional.of(new CleanupConfiguration(
+ annotation.truncateTables(),
+ List.copyOf(Arrays.asList(annotation.excludeTruncateTables()))));
+ }
+ return findSingletonAnnotation(testClass)
+ .map(annotation -> new CleanupConfiguration(
+ annotation.truncateTables(),
+ List.copyOf(Arrays.asList(annotation.excludeTruncateTables()))));
}
- private static Optional findSingletonAnnotation(ExtensionContext context) {
- return AnnotationSupport.findAnnotation(context.getTestClass(), PostgresqlTestcontainerSingleton.class);
+ private static Optional findPrototypeAnnotation(Class> testClass) {
+ return AnnotationSupport.findAnnotation(testClass, PostgresqlTestcontainer.class);
}
private static Optional findSingletonAnnotation(Class> testClass) {
return AnnotationSupport.findAnnotation(testClass, PostgresqlTestcontainerSingleton.class);
}
+ private static final class ContainerReference {
+
+ private final PostgresqlContainerExtension container;
+ private final boolean singleton;
+
+ private ContainerReference(PostgresqlContainerExtension container, boolean singleton) {
+ this.container = container;
+ this.singleton = singleton;
+ }
+
+ private PostgresqlContainerExtension container() {
+ return container;
+ }
+
+ private boolean singleton() {
+ return singleton;
+ }
+
+ }
+
+ private record CleanupConfiguration(boolean truncateTables, List excludedTables) {
+ }
+
public static class PostgresqlTestcontainerContextCustomizerFactory implements ContextCustomizerFactory {
@Override
public ContextCustomizer createContextCustomizer(
Class> testClass,
List configAttributes) {
- return (context, mergedConfig) -> {
- if (findPrototypeAnnotation(testClass).isPresent()) {
- init(context, findPrototypeAnnotation(testClass).get().properties());
- } else if (findSingletonAnnotation(testClass).isPresent()) {
- init(context, findSingletonAnnotation(testClass).get().properties());
- }
- };
+ var prototype = findPrototypeAnnotation(testClass);
+ if (prototype.isPresent()) {
+ return new PostgresqlContextCustomizer(
+ testClass,
+ false,
+ List.copyOf(Arrays.asList(prototype.get().properties())));
+ }
+ var singleton = findSingletonAnnotation(testClass);
+ if (singleton.isPresent()) {
+ return new PostgresqlContextCustomizer(
+ testClass,
+ true,
+ List.copyOf(Arrays.asList(singleton.get().properties())));
+ }
+ return null;
}
+ }
- private void init(ConfigurableApplicationContext context, String[] properties) {
- var container = THREAD_CONTAINER.get();
+ private record PostgresqlContextCustomizer(
+ Class> testClass,
+ boolean singleton,
+ List properties) implements ContextCustomizer {
+
+ @Override
+ public void customizeContext(
+ ConfigurableApplicationContext context,
+ MergedContextConfiguration mergedConfig) {
+ var container = getOrStart(testClass).container();
var jdbcUrl = container.getJdbcUrl();
var username = container.getUsername();
var password = container.getPassword();
@@ -147,7 +208,7 @@ private void init(ConfigurableApplicationContext context, String[] properties) {
"flyway.user=" + username,
"flyway.password=" + password,
"flyway.postgresql.transactional.lock=false")
- .and(properties)
+ .and(properties.toArray(String[]::new))
.applyTo(context);
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlTestcontainerFactory.java b/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlTestcontainerFactory.java
index c3df4630..6192658f 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlTestcontainerFactory.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/postgresql/PostgresqlTestcontainerFactory.java
@@ -1,49 +1,60 @@
package dev.vality.testcontainers.annotations.postgresql;
+import dev.vality.testcontainers.annotations.util.ContainerShutdownRegistry;
import lombok.AccessLevel;
import lombok.NoArgsConstructor;
-import lombok.Synchronized;
/**
- * Фабрика по созданию контейнеров
- * {@link #create()} создает экземпляр тестконтейнера
- *
{@link #getOrCreateSingletonContainer()} создает синглтон тестконтейнера
- *
- * @see PostgresqlTestcontainerExtension PostgresqlTestcontainerExtension
+ * Фабрика по созданию PostgreSQL контейнеров.
*/
@NoArgsConstructor(access = AccessLevel.PRIVATE)
public class PostgresqlTestcontainerFactory {
- private PostgresqlContainerExtension postgresqlContainer;
+ private volatile PostgresqlContainerExtension singletonContainer;
public static PostgresqlContainerExtension container() {
- return instance().create();
+ return new PostgresqlContainerExtension();
}
public static PostgresqlContainerExtension singletonContainer() {
return instance().getOrCreateSingletonContainer();
}
+ static void discardSingleton(PostgresqlContainerExtension container) {
+ instance().discard(container);
+ }
+
private static PostgresqlTestcontainerFactory instance() {
return SingletonHolder.INSTANCE;
}
- @Synchronized
- private PostgresqlContainerExtension getOrCreateSingletonContainer() {
- if (postgresqlContainer != null) {
- return postgresqlContainer;
+ private void discard(PostgresqlContainerExtension container) {
+ synchronized (this) {
+ if (singletonContainer != container) {
+ return;
+ }
+ singletonContainer = null;
}
- postgresqlContainer = create();
- return postgresqlContainer;
+ ContainerShutdownRegistry.unregister(container);
+ container.stop();
}
- private PostgresqlContainerExtension create() {
- return new PostgresqlContainerExtension();
+ private PostgresqlContainerExtension getOrCreateSingletonContainer() {
+ var result = singletonContainer;
+ if (result == null) {
+ synchronized (this) {
+ result = singletonContainer;
+ if (result == null) {
+ result = ContainerShutdownRegistry.register(new PostgresqlContainerExtension());
+ singletonContainer = result;
+ }
+ }
+ }
+ return result;
}
private static class SingletonHolder {
private static final PostgresqlTestcontainerFactory INSTANCE = new PostgresqlTestcontainerFactory();
-
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/util/ContainerShutdownRegistry.java b/src/main/java/dev/vality/testcontainers/annotations/util/ContainerShutdownRegistry.java
new file mode 100644
index 00000000..e9ff5084
--- /dev/null
+++ b/src/main/java/dev/vality/testcontainers/annotations/util/ContainerShutdownRegistry.java
@@ -0,0 +1,42 @@
+package dev.vality.testcontainers.annotations.util;
+
+import lombok.extern.slf4j.Slf4j;
+import org.testcontainers.lifecycle.Startable;
+
+import java.util.Set;
+import java.util.concurrent.ConcurrentHashMap;
+
+@Slf4j
+public final class ContainerShutdownRegistry {
+
+ private static final Set CONTAINERS = ConcurrentHashMap.newKeySet();
+
+ static {
+ Runtime.getRuntime().addShutdownHook(new Thread(
+ ContainerShutdownRegistry::stopAll,
+ "testcontainers-annotations-shutdown"));
+ }
+
+ private ContainerShutdownRegistry() {
+ }
+
+ public static T register(T container) {
+ CONTAINERS.add(container);
+ return container;
+ }
+
+ public static void unregister(Startable container) {
+ CONTAINERS.remove(container);
+ }
+
+ private static void stopAll() {
+ CONTAINERS.forEach(container -> {
+ try {
+ container.stop();
+ } catch (RuntimeException ex) {
+ log.warn("Unable to stop shared test container", ex);
+ }
+ });
+ CONTAINERS.clear();
+ }
+}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/util/GenericContainerUtil.java b/src/main/java/dev/vality/testcontainers/annotations/util/GenericContainerUtil.java
index 13e20325..61bb6a00 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/util/GenericContainerUtil.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/util/GenericContainerUtil.java
@@ -1,30 +1,41 @@
package dev.vality.testcontainers.annotations.util;
import dev.vality.testcontainers.annotations.kafka.KafkaContainerExtension;
+import org.testcontainers.containers.ContainerState;
import org.testcontainers.containers.GenericContainer;
import org.testcontainers.containers.wait.strategy.HttpWaitStrategy;
import org.testcontainers.containers.wait.strategy.WaitStrategy;
+import org.testcontainers.lifecycle.Startable;
import org.testcontainers.lifecycle.Startables;
import java.time.Duration;
+import java.util.concurrent.CompletionException;
import java.util.stream.Stream;
-import static org.assertj.core.api.Assertions.assertThat;
-
public class GenericContainerUtil {
public static void startContainer(GenericContainer> container) {
- Startables.deepStart(Stream.of(container))
- .join();
- assertThat(container.isRunning())
- .isTrue();
+ start(container);
}
public static void startContainer(KafkaContainerExtension container) {
- Startables.deepStart(Stream.of(container))
- .join();
- assertThat(container.isRunning())
- .isTrue();
+ start(container);
+ }
+
+ private static void start(T container) {
+ synchronized (container) {
+ if (container.isRunning()) {
+ return;
+ }
+ try {
+ Startables.deepStart(Stream.of(container)).join();
+ } catch (CompletionException ex) {
+ throw new IllegalStateException("Unable to start test container", ex.getCause());
+ }
+ if (!container.isRunning()) {
+ throw new IllegalStateException("Test container did not reach the running state");
+ }
+ }
}
public static WaitStrategy getWaitStrategy(String path, Integer statusCode, Integer port, Duration duration) {
diff --git a/src/main/java/dev/vality/testcontainers/annotations/util/RandomBeans.java b/src/main/java/dev/vality/testcontainers/annotations/util/RandomBeans.java
index 2111218d..aabdc585 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/util/RandomBeans.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/util/RandomBeans.java
@@ -15,98 +15,99 @@
import java.util.Calendar;
import java.util.Date;
import java.util.List;
+import java.util.TimeZone;
+import java.util.concurrent.ThreadLocalRandom;
import java.util.stream.Collectors;
import java.util.stream.Stream;
public class RandomBeans {
public static final long DEFAULT_SEED = 123L;
+ private static final Instant BASE_INSTANT = Instant.parse("2020-01-01T00:00:00Z");
+
public static T random(Class type, String... excludedFields) {
- var parameters = createParametersWithExcludedFields(DEFAULT_SEED, excludedFields);
- var easyRandom = new EasyRandom(parameters);
+ var easyRandom = new EasyRandom(createParametersWithExcludedFields(
+ randomSeed(), Clock.systemUTC(), excludedFields));
return easyRandom.nextObject(type);
}
public static T random(Long seed, Class type, String... excludedFields) {
- var parameters = createParametersWithExcludedFields(seed, excludedFields);
- var easyRandom = new EasyRandom(parameters);
+ var easyRandom = new EasyRandom(createDeterministicParameters(seed, excludedFields));
return easyRandom.nextObject(type);
}
public static List randomListOf(int amount, Class type, String... excludedFields) {
- var parameters = createParametersWithExcludedFields(DEFAULT_SEED, excludedFields);
- var easyRandom = new EasyRandom(parameters);
+ var easyRandom = new EasyRandom(createParametersWithExcludedFields(
+ randomSeed(), Clock.systemUTC(), excludedFields));
return easyRandom.objects(type, amount).collect(Collectors.toList());
}
public static List randomListOf(Long seed, int amount, Class type, String... excludedFields) {
- var parameters = createParametersWithExcludedFields(seed, excludedFields);
- var easyRandom = new EasyRandom(parameters);
+ var easyRandom = new EasyRandom(createDeterministicParameters(seed, excludedFields));
return easyRandom.objects(type, amount).collect(Collectors.toList());
}
public static Stream randomStreamOf(int amount, Class type, String... excludedFields) {
- var parameters = createParametersWithExcludedFields(DEFAULT_SEED, excludedFields);
- var easyRandom = new EasyRandom(parameters);
+ var easyRandom = new EasyRandom(createParametersWithExcludedFields(
+ randomSeed(), Clock.systemUTC(), excludedFields));
return easyRandom.objects(type, amount);
}
@SneakyThrows
public static > T randomThrift(Class type) {
- var mockTBaseProcessor = new MockTBaseProcessor(MockMode.ALL, 25, 1);
+ return randomThrift(type, Clock.systemUTC(), MockMode.ALL);
+ }
+
+ @SneakyThrows
+ public static > T randomThrift(Class type, Clock clock) {
+ return randomThrift(type, clock, MockMode.ALL);
+ }
+
+ @SneakyThrows
+ private static > T randomThrift(
+ Class type,
+ Clock clock,
+ MockMode mode) {
+ var mockTBaseProcessor = new MockTBaseProcessor(mode, 25, 1);
mockTBaseProcessor.addFieldHandler(
- structHandler -> structHandler.value(Instant.now().toString()),
+ structHandler -> structHandler.value(Instant.now(clock).toString()),
"created_at", "at", "due");
return mockTBaseProcessor.process(type.getConstructor().newInstance(), new TBaseHandler<>(type));
}
@SneakyThrows
public static > T randomThriftOnlyRequiredFields(Class type) {
- var mockTBaseProcessor = new MockTBaseProcessor(MockMode.REQUIRED_ONLY, 25, 1);
- mockTBaseProcessor.addFieldHandler(
- structHandler -> structHandler.value(Instant.now().toString()),
- "created_at", "at", "due");
- return mockTBaseProcessor.process(type.getConstructor().newInstance(), new TBaseHandler<>(type));
+ return randomThrift(type, Clock.systemUTC(), MockMode.REQUIRED_ONLY);
+ }
+
+ private static EasyRandomParameters createDeterministicParameters(Long seed, String... excludedFields) {
+ if (seed == null) {
+ throw new IllegalArgumentException("Seed must not be null");
+ }
+ return createParametersWithExcludedFields(seed, deterministicClock(seed), excludedFields);
}
- private static EasyRandomParameters createParametersWithExcludedFields(Long seed, String... excludedFields) {
+ private static EasyRandomParameters createParametersWithExcludedFields(
+ Long seed,
+ Clock clock,
+ String... excludedFields) {
+ var instant = Instant.now(clock).truncatedTo(ChronoUnit.MICROS);
+ var localDateTime = LocalDateTime.ofInstant(instant, ZoneOffset.UTC);
var parameters = new EasyRandomParameters();
- parameters.randomize(LocalDateTime.class, () -> {
- var dateTime = LocalDateTime.now();
- return dateTime.truncatedTo(ChronoUnit.MICROS);
- });
- parameters.randomize(Instant.class, () -> {
- var instant = Instant.now();
- return instant.truncatedTo(ChronoUnit.MICROS);
- });
- parameters.randomize(Date.class, () -> {
- var instant = Instant.now().truncatedTo(ChronoUnit.MICROS);
- return Date.from(instant);
- });
- parameters.randomize(Timestamp.class, () -> {
- var instant = Instant.now().truncatedTo(ChronoUnit.MICROS);
- return Timestamp.from(instant);
- });
- parameters.randomize(OffsetDateTime.class, () -> {
- var offsetDateTime = OffsetDateTime.now();
- return offsetDateTime.truncatedTo(ChronoUnit.MICROS);
- });
- parameters.randomize(ZonedDateTime.class, () -> {
- var zonedDateTime = ZonedDateTime.now();
- return zonedDateTime.truncatedTo(ChronoUnit.MICROS);
- });
+ parameters.randomize(LocalDateTime.class, () -> localDateTime);
+ parameters.randomize(Instant.class, () -> instant);
+ parameters.randomize(Date.class, () -> Date.from(instant));
+ parameters.randomize(Timestamp.class, () -> Timestamp.from(instant));
+ parameters.randomize(OffsetDateTime.class, () -> OffsetDateTime.ofInstant(instant, ZoneOffset.UTC));
+ parameters.randomize(ZonedDateTime.class, () -> ZonedDateTime.ofInstant(instant, ZoneOffset.UTC));
parameters.randomize(Calendar.class, () -> {
- var instant = Instant.now().truncatedTo(ChronoUnit.MICROS);
- var calendar = Calendar.getInstance();
+ var calendar = Calendar.getInstance(TimeZone.getTimeZone("UTC"));
calendar.setTime(Date.from(instant));
return calendar;
});
- parameters.randomize(LocalDate.class, LocalDate::now);
- parameters.randomize(LocalTime.class, () -> {
- var time = LocalTime.now();
- return time.truncatedTo(ChronoUnit.MICROS);
- });
+ parameters.randomize(LocalDate.class, localDateTime::toLocalDate);
+ parameters.randomize(LocalTime.class, localDateTime::toLocalTime);
if (excludedFields != null) {
for (var excludedField : excludedFields) {
parameters.excludeField(field -> field.getName().equals(excludedField));
@@ -120,4 +121,14 @@ private static EasyRandomParameters createParametersWithExcludedFields(Long seed
.collectionSizeRange(1, 10);
return parameters;
}
+
+ private static Clock deterministicClock(long seed) {
+ var secondsInTenYears = 3650L * 24 * 60 * 60;
+ var offsetSeconds = Math.floorMod(seed, secondsInTenYears);
+ return Clock.fixed(BASE_INSTANT.plusSeconds(offsetSeconds), ZoneOffset.UTC);
+ }
+
+ private static long randomSeed() {
+ return ThreadLocalRandom.current().nextLong();
+ }
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/util/SharedTestResourceLock.java b/src/main/java/dev/vality/testcontainers/annotations/util/SharedTestResourceLock.java
new file mode 100644
index 00000000..ae11d60f
--- /dev/null
+++ b/src/main/java/dev/vality/testcontainers/annotations/util/SharedTestResourceLock.java
@@ -0,0 +1,70 @@
+package dev.vality.testcontainers.annotations.util;
+
+import lombok.AccessLevel;
+import lombok.NoArgsConstructor;
+
+import java.util.concurrent.*;
+import java.util.concurrent.atomic.AtomicBoolean;
+import java.util.concurrent.atomic.AtomicInteger;
+
+/**
+ * Serializes test classes that use process-wide singleton containers.
+ * A semaphore is used intentionally because JUnit may invoke lifecycle callbacks on different threads.
+ */
+@NoArgsConstructor(access = AccessLevel.PRIVATE)
+public final class SharedTestResourceLock {
+
+ private static final Semaphore SINGLETON_TESTS = new Semaphore(1, true);
+ private static final ConcurrentMap, Holder> HOLDERS = new ConcurrentHashMap<>();
+
+ public static void acquire(Class> testClass) {
+ var owner = new AtomicBoolean();
+ var holder = HOLDERS.compute(testClass, (ignored, existing) -> {
+ if (existing == null) {
+ owner.set(true);
+ return new Holder();
+ }
+ existing.references.incrementAndGet();
+ return existing;
+ });
+
+ if (owner.get()) {
+ try {
+ SINGLETON_TESTS.acquire();
+ holder.acquired.complete(null);
+ } catch (InterruptedException ex) {
+ Thread.currentThread().interrupt();
+ HOLDERS.remove(testClass, holder);
+ holder.acquired.completeExceptionally(ex);
+ throw new IllegalStateException("Interrupted while waiting for singleton test resources", ex);
+ }
+ return;
+ }
+
+ try {
+ holder.acquired.join();
+ } catch (CompletionException ex) {
+ throw new IllegalStateException("Unable to acquire singleton test resources", ex.getCause());
+ }
+ }
+
+ public static void release(Class> testClass) {
+ var releaseSemaphore = new AtomicBoolean();
+ HOLDERS.computeIfPresent(testClass, (ignored, holder) -> {
+ if (holder.references.decrementAndGet() == 0) {
+ releaseSemaphore.set(true);
+ return null;
+ }
+ return holder;
+ });
+ if (releaseSemaphore.get()) {
+ SINGLETON_TESTS.release();
+ }
+ }
+
+ private static final class Holder {
+
+ private final AtomicInteger references = new AtomicInteger(1);
+ private final CompletableFuture acquired = new CompletableFuture<>();
+ }
+}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/util/SpringApplicationPropertiesLoader.java b/src/main/java/dev/vality/testcontainers/annotations/util/SpringApplicationPropertiesLoader.java
index 9f29d00a..17bdfbb8 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/util/SpringApplicationPropertiesLoader.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/util/SpringApplicationPropertiesLoader.java
@@ -2,108 +2,198 @@
import dev.vality.testcontainers.annotations.exception.IoException;
import dev.vality.testcontainers.annotations.exception.NoSuchFileException;
-import lombok.Builder;
-import lombok.Data;
import org.springframework.boot.env.PropertiesPropertySourceLoader;
import org.springframework.boot.env.PropertySourceLoader;
import org.springframework.boot.env.YamlPropertySourceLoader;
import org.springframework.boot.origin.OriginTrackedValue;
+import org.springframework.core.env.EnumerablePropertySource;
+import org.springframework.core.env.Profiles;
+import org.springframework.core.env.PropertySource;
import org.springframework.core.io.ClassPathResource;
import java.io.IOException;
-import java.util.AbstractMap;
-import java.util.List;
-import java.util.Map;
-import java.util.Properties;
+import java.util.*;
import java.util.function.Supplier;
-import java.util.stream.Collectors;
-
-import static org.assertj.core.api.Assertions.assertThat;
public class SpringApplicationPropertiesLoader {
- private static final List>> TYPES = List.of(
- Map.entry("yml", YamlPropertySourceLoader::new),
- Map.entry("yaml", YamlPropertySourceLoader::new),
- Map.entry("properties", PropertiesPropertySourceLoader::new),
- Map.entry("xml", PropertiesPropertySourceLoader::new)
- );
+ private static final List TYPES = List.of(
+ new FileType("yml", YamlPropertySourceLoader::new),
+ new FileType("yaml", YamlPropertySourceLoader::new),
+ new FileType("properties", PropertiesPropertySourceLoader::new),
+ new FileType("xml", PropertiesPropertySourceLoader::new));
public static String loadDefaultLibraryProperty(String key) {
- Object tag;
- try {
- tag = loadPropertiesByFile().get(key);
- } catch (NoSuchFileException ex) {
- tag = null;
+ return findExternalValue(key)
+ .or(() -> findValue(loadApplicationProperties(false), key))
+ .or(() -> findValue(loadNamedProperties("testcontainers-annotations", true), key))
+ .map(String::valueOf)
+ .filter(value -> !value.isBlank())
+ .orElseThrow(() -> new IllegalStateException("Required property is not configured: " + key));
+ }
+
+ public static Properties loadFromSpringApplicationPropertiesFile(List keys) {
+ if (keys == null || keys.isEmpty()) {
+ return new Properties();
}
- if (tag == null) {
- tag = getSource(findPropertiesFileParametersByName("testcontainers-annotations")).get(key);
+ var applicationProperties = loadApplicationProperties(true);
+ var result = new Properties();
+ var missing = new ArrayList();
+ for (var key : keys) {
+ var value = findExternalValue(key)
+ .or(() -> findValue(applicationProperties, key));
+ if (value.isPresent()) {
+ result.setProperty(key, String.valueOf(value.get()));
+ } else {
+ missing.add(key);
+ }
}
- return String.valueOf(tag);
+ if (!missing.isEmpty()) {
+ throw new IllegalStateException("Required application properties are missing: " + missing);
+ }
+ return result;
}
- public static Properties loadFromSpringApplicationPropertiesFile(List keys) {
- var fileProperties = loadPropertiesByFile();
- var filtered = fileProperties.entrySet().stream()
- .filter(entry -> keys.contains(entry.getKey()))
- .collect(Collectors.toMap(Map.Entry::getKey, Map.Entry::getValue));
- assertThat(filtered.keySet()).containsAll(keys);
- var properties = new Properties();
- properties.putAll(filtered);
+ private static Map loadApplicationProperties(boolean required) {
+ var baseDocuments = loadNamedPropertyDocuments("application", required);
+ var unconditionalProperties = mergeDocuments(baseDocuments, List.of(), false);
+ var profiles = activeProfiles(unconditionalProperties);
+
+ var properties = new LinkedHashMap<>(mergeDocuments(baseDocuments, profiles, true));
+ profiles.forEach(profile -> properties.putAll(mergeDocuments(
+ loadNamedPropertyDocuments("application-" + profile, false),
+ profiles,
+ true)));
return properties;
}
- private static Map loadPropertiesByFile() {
- var parameters = findPropertiesFileParameters();
- return getSource(parameters);
+ private static Map loadNamedProperties(String baseName, boolean required) {
+ var result = new LinkedHashMap();
+ loadNamedPropertyDocuments(baseName, required).forEach(result::putAll);
+ return result;
}
- private static Map getSource(PropertiesFileParameters parameters) {
- var currentClass = SpringApplicationPropertiesLoader.class;
- var classPathResource = new ClassPathResource(parameters.getName(), currentClass.getClassLoader());
- try {
- //noinspection unchecked
- return ((Map) parameters.getPropertySourceLoader().get()
- .load(classPathResource.getFilename(), classPathResource)
- .getFirst()
- .getSource())
- .entrySet().stream()
- .map(entry -> new AbstractMap.SimpleEntry<>(entry.getKey(), entry.getValue().getValue()))
- .collect(Collectors.toMap(Map.Entry::getKey, Map.Entry::getValue));
- } catch (IOException ex) {
- throw new IoException("Error when loading properties, ", ex);
+ private static List> loadNamedPropertyDocuments(String baseName, boolean required) {
+ var documents = new ArrayList>();
+ var found = false;
+ for (var type : TYPES) {
+ var resource = new ClassPathResource(baseName + "." + type.extension());
+ if (!resource.exists()) {
+ continue;
+ }
+ found = true;
+ try {
+ var loader = type.loader().get();
+ var propertySources = loader.load(resource.getFilename(), resource);
+ for (var propertySource : propertySources) {
+ var document = new LinkedHashMap();
+ mergePropertySource(document, propertySource);
+ documents.add(document);
+ }
+ } catch (IOException ex) {
+ throw new IoException("Error while loading " + resource.getPath(), ex);
+ }
}
+ if (required && !found) {
+ throw new NoSuchFileException(
+ "Configuration file not found: " + baseName + ".[yml|yaml|properties|xml]");
+ }
+ return documents;
}
- private static PropertiesFileParameters findPropertiesFileParameters() {
- return findPropertiesFileParametersByName("application");
+ private static Map mergeDocuments(
+ List> documents,
+ List activeProfiles,
+ boolean includeActiveDocuments) {
+ var result = new LinkedHashMap();
+ for (var document : documents) {
+ var activation = profileActivation(document);
+ if (activation.isEmpty() || includeActiveDocuments && matchesProfiles(activation.get(), activeProfiles)) {
+ result.putAll(document);
+ }
+ }
+ return result;
}
- private static PropertiesFileParameters findPropertiesFileParametersByName(String name) {
- var currentClass = SpringApplicationPropertiesLoader.class;
- return TYPES.stream()
- .map(entry ->
- Map.entry("%s.%s".formatted(name, entry.getKey()), entry.getValue())
- )
- .filter(entry -> currentClass.getResource("/" + entry.getKey()) != null)
- .findFirst()
- .map(entry -> PropertiesFileParameters.builder()
- .propertySourceLoader(entry.getValue())
- .name(entry.getKey())
- .build())
- .orElseThrow(() -> new NoSuchFileException(
- "Error loading configuration: " +
- "src/main/resources/application.[yml|yaml|properties|xml] — " +
- "file not found"
- ));
+ private static Optional profileActivation(Map document) {
+ return findValue(document, "spring.config.activate.on-profile")
+ .or(() -> findValue(document, "spring.profiles"))
+ .map(String::valueOf)
+ .map(String::trim)
+ .filter(value -> !value.isEmpty());
}
- @Data
- @Builder
- private static class PropertiesFileParameters {
+ private static boolean matchesProfiles(String expression, List activeProfiles) {
+ var expressions = Arrays.stream(expression.split(","))
+ .map(String::trim)
+ .filter(value -> !value.isEmpty())
+ .toArray(String[]::new);
+ return expressions.length > 0 && Profiles.of(expressions).matches(activeProfiles::contains);
+ }
- private Supplier propertySourceLoader;
- private String name;
+ private static void mergePropertySource(Map target, PropertySource> propertySource) {
+ if (propertySource instanceof EnumerablePropertySource> enumerable) {
+ for (var propertyName : enumerable.getPropertyNames()) {
+ var value = enumerable.getProperty(propertyName);
+ if (value != null) {
+ target.put(propertyName, unwrap(value));
+ }
+ }
+ return;
+ }
+ if (propertySource.getSource() instanceof Map, ?> source) {
+ source.forEach((key, value) -> {
+ if (key != null && value != null) {
+ target.put(String.valueOf(key), unwrap(value));
+ }
+ });
+ }
+ }
+
+ private static Object unwrap(Object value) {
+ return value instanceof OriginTrackedValue originTrackedValue
+ ? originTrackedValue.getValue()
+ : value;
+ }
+
+ private static List activeProfiles(Map applicationProperties) {
+ var profiles = new LinkedHashSet();
+ addProfiles(
+ profiles,
+ findExternalValue("spring.profiles.active")
+ .or(() -> findValue(applicationProperties, "spring.profiles.active")));
+ addProfiles(
+ profiles,
+ findExternalValue("spring.profiles.include")
+ .or(() -> findValue(applicationProperties, "spring.profiles.include")));
+ if (profiles.isEmpty()) {
+ profiles.add("default");
+ }
+ return List.copyOf(profiles);
+ }
+
+ private static void addProfiles(Set target, Optional value) {
+ value.map(String::valueOf)
+ .stream()
+ .flatMap(profiles -> Arrays.stream(profiles.split(",")))
+ .map(String::trim)
+ .filter(profile -> !profile.isEmpty())
+ .forEach(target::add);
+ }
+
+ private static Optional findValue(Map source, String key) {
+ return Optional.ofNullable(source.get(key));
+ }
+
+ private static Optional findExternalValue(String key) {
+ var systemProperty = System.getProperty(key);
+ if (systemProperty != null) {
+ return Optional.of(systemProperty);
+ }
+ var environmentKey = key.toUpperCase(Locale.ROOT).replace('.', '_').replace('-', '_');
+ return Optional.ofNullable(System.getenv(environmentKey));
+ }
+ private record FileType(String extension, Supplier loader) {
}
}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/util/TestExecutionLock.java b/src/main/java/dev/vality/testcontainers/annotations/util/TestExecutionLock.java
new file mode 100644
index 00000000..1cb89279
--- /dev/null
+++ b/src/main/java/dev/vality/testcontainers/annotations/util/TestExecutionLock.java
@@ -0,0 +1,63 @@
+package dev.vality.testcontainers.annotations.util;
+
+import lombok.AccessLevel;
+import lombok.NoArgsConstructor;
+import org.junit.jupiter.api.extension.ExtensionContext;
+
+import java.util.concurrent.Semaphore;
+import java.util.concurrent.atomic.AtomicBoolean;
+
+/**
+ * Serializes test methods that mutate the same class-scoped test resources.
+ * The lock is released automatically when the JUnit method context is closed.
+ */
+@NoArgsConstructor(access = AccessLevel.PRIVATE)
+public final class TestExecutionLock {
+
+ private static final ExtensionContext.Namespace NAMESPACE =
+ ExtensionContext.Namespace.create(TestExecutionLock.class);
+ private static final String METHOD_LOCK = "method-lock";
+
+ public static void acquire(ExtensionContext context) {
+ var methodStore = context.getStore(NAMESPACE);
+ if (methodStore.get(METHOD_LOCK) != null) {
+ return;
+ }
+
+ var semaphore = context.getRoot().getStore(NAMESPACE).getOrComputeIfAbsent(
+ context.getRequiredTestClass(),
+ ignored -> new Semaphore(1, true),
+ Semaphore.class);
+ try {
+ semaphore.acquire();
+ } catch (InterruptedException ex) {
+ Thread.currentThread().interrupt();
+ throw new IllegalStateException("Interrupted while waiting for test resource", ex);
+ }
+ methodStore.put(METHOD_LOCK, new LockHandle(semaphore));
+ }
+
+ public static void release(ExtensionContext context) {
+ var handle = context.getStore(NAMESPACE).remove(METHOD_LOCK, LockHandle.class);
+ if (handle != null) {
+ handle.close();
+ }
+ }
+
+ private static final class LockHandle implements ExtensionContext.Store.CloseableResource {
+
+ private final Semaphore semaphore;
+ private final AtomicBoolean closed = new AtomicBoolean();
+
+ private LockHandle(Semaphore semaphore) {
+ this.semaphore = semaphore;
+ }
+
+ @Override
+ public void close() {
+ if (closed.compareAndSet(false, true)) {
+ semaphore.release();
+ }
+ }
+ }
+}
diff --git a/src/main/java/dev/vality/testcontainers/annotations/util/ValuesGenerator.java b/src/main/java/dev/vality/testcontainers/annotations/util/ValuesGenerator.java
index bab62029..27b4770f 100644
--- a/src/main/java/dev/vality/testcontainers/annotations/util/ValuesGenerator.java
+++ b/src/main/java/dev/vality/testcontainers/annotations/util/ValuesGenerator.java
@@ -8,77 +8,89 @@
import java.io.IOException;
import java.io.InputStream;
import java.nio.charset.StandardCharsets;
-import java.time.Instant;
-import java.time.LocalDateTime;
-import java.time.ZoneId;
-import java.time.ZoneOffset;
+import java.time.*;
import java.time.temporal.ChronoUnit;
import java.util.UUID;
+import java.util.concurrent.atomic.AtomicLong;
@NoArgsConstructor(access = AccessLevel.PRIVATE)
public class ValuesGenerator {
- private static final LocalDateTime fromTime = LocalDateTime.now().minusHours(3);
- private static final LocalDateTime toTime = LocalDateTime.now().minusHours(1);
- private static final LocalDateTime inFromToPeriodTime = LocalDateTime.now().minusHours(2);
+ private static final AtomicLong SEED_SEQUENCE = new AtomicLong(RandomBeans.DEFAULT_SEED);
+ private static volatile Clock clock = Clock.systemDefaultZone();
public static String generateId() {
return UUID.randomUUID().toString();
}
public static String generateDate() {
- return TypeUtil.temporalToString(LocalDateTime.now().truncatedTo(ChronoUnit.MICROS));
+ return TypeUtil.temporalToString(LocalDateTime.now(clock).truncatedTo(ChronoUnit.MICROS));
}
public static Long generateLong() {
- return RandomBeans.random(Long.class);
+ return RandomBeans.random(SEED_SEQUENCE.getAndIncrement(), Long.class);
}
public static Integer generateInt() {
- return RandomBeans.random(Integer.class);
+ return RandomBeans.random(SEED_SEQUENCE.getAndIncrement(), Integer.class);
}
public static String generateString() {
- return RandomBeans.random(String.class);
+ return RandomBeans.random(SEED_SEQUENCE.getAndIncrement(), String.class);
}
public static LocalDateTime generateLocalDateTime() {
- return RandomBeans.random(LocalDateTime.class).truncatedTo(ChronoUnit.MICROS);
+ return RandomBeans.random(SEED_SEQUENCE.getAndIncrement(), LocalDateTime.class)
+ .truncatedTo(ChronoUnit.MICROS);
}
public static Instant generateInstant() {
- return RandomBeans.random(Instant.class).truncatedTo(ChronoUnit.MICROS);
+ return RandomBeans.random(SEED_SEQUENCE.getAndIncrement(), Instant.class)
+ .truncatedTo(ChronoUnit.MICROS);
}
public static Instant generateCurrentTimePlusDay() {
- return LocalDateTime.now().plusDays(1).toInstant(getZoneOffset()).truncatedTo(ChronoUnit.MICROS);
+ return ZonedDateTime.now(clock).plusDays(1).toInstant().truncatedTo(ChronoUnit.MICROS);
}
public static Instant generateCurrentTimePlusSecond() {
- return LocalDateTime.now().plusSeconds(1).toInstant(getZoneOffset()).truncatedTo(ChronoUnit.MICROS);
+ return ZonedDateTime.now(clock).plusSeconds(1).toInstant().truncatedTo(ChronoUnit.MICROS);
}
public static ZoneOffset getZoneOffset() {
- return ZoneId.systemDefault().getRules().getOffset(LocalDateTime.now());
+ return ZonedDateTime.now(clock).getOffset();
}
public static String getContent(InputStream content) throws IOException {
- return IOUtils.toString(content, StandardCharsets.UTF_8);
+ try (content) {
+ return IOUtils.toString(content, StandardCharsets.UTF_8);
+ }
}
public static LocalDateTime getFromTime() {
- return fromTime.truncatedTo(ChronoUnit.MICROS);
+ return LocalDateTime.now(clock).minusHours(3).truncatedTo(ChronoUnit.MICROS);
}
public static LocalDateTime getToTime() {
- return toTime.truncatedTo(ChronoUnit.MICROS);
+ return LocalDateTime.now(clock).minusHours(1).truncatedTo(ChronoUnit.MICROS);
}
public static LocalDateTime getInFromToPeriodTime() {
- return inFromToPeriodTime.truncatedTo(ChronoUnit.MICROS);
+ return LocalDateTime.now(clock).minusHours(2).truncatedTo(ChronoUnit.MICROS);
}
public static Instant getCurrentInstant() {
- return Instant.now().truncatedTo(ChronoUnit.MICROS);
+ return Instant.now(clock).truncatedTo(ChronoUnit.MICROS);
+ }
+
+ public static void useClock(Clock customClock) {
+ if (customClock == null) {
+ throw new IllegalArgumentException("Clock must not be null");
+ }
+ clock = customClock;
+ }
+
+ public static void resetClock() {
+ clock = Clock.systemDefaultZone();
}
}
diff --git a/src/test/java/dev/vality/testcontainers/annotations/clickhouse/SqlScriptParserTest.java b/src/test/java/dev/vality/testcontainers/annotations/clickhouse/SqlScriptParserTest.java
new file mode 100644
index 00000000..fa3e83e0
--- /dev/null
+++ b/src/test/java/dev/vality/testcontainers/annotations/clickhouse/SqlScriptParserTest.java
@@ -0,0 +1,24 @@
+package dev.vality.testcontainers.annotations.clickhouse;
+
+import org.junit.jupiter.api.Test;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+class SqlScriptParserTest {
+
+ @Test
+ void shouldIgnoreSemicolonsInsideQuotedValuesAndComments() {
+ var statements = SqlScriptParser.splitStatements("""
+ CREATE TABLE events (value String);
+ INSERT INTO events VALUES ('a;b'); -- comment ;
+ /* block ; comment */ SELECT $$dollar;quoted$$;
+ """);
+
+ assertThat(statements)
+ .hasSize(3)
+ .element(1)
+ .isEqualTo("INSERT INTO events VALUES ('a;b')");
+ assertThat(statements.get(2))
+ .contains("$$dollar;quoted$$");
+ }
+}
diff --git a/src/test/java/dev/vality/testcontainers/annotations/util/RandomBeansTest.java b/src/test/java/dev/vality/testcontainers/annotations/util/RandomBeansTest.java
new file mode 100644
index 00000000..939ac088
--- /dev/null
+++ b/src/test/java/dev/vality/testcontainers/annotations/util/RandomBeansTest.java
@@ -0,0 +1,36 @@
+package dev.vality.testcontainers.annotations.util;
+
+import lombok.Data;
+import org.junit.jupiter.api.Test;
+
+import java.time.Instant;
+
+import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.junit.jupiter.api.Assertions.assertNotEquals;
+
+class RandomBeansTest {
+
+ @Test
+ void shouldKeepSeededGenerationDeterministic() {
+ var first = RandomBeans.random(123L, RandomBean.class);
+ var second = RandomBeans.random(123L, RandomBean.class);
+
+ assertEquals(first, second);
+ }
+
+ @Test
+ void shouldUseCurrentTimeForGenerationWithoutExplicitSeed() {
+ var first = RandomBeans.random(RandomBean.class);
+ var second = RandomBeans.random(RandomBean.class);
+
+ assertNotEquals(first.getValue(), second.getValue());
+ assertNotEquals(first.getCreatedAt(), second.getCreatedAt());
+ }
+
+ @Data
+ public static class RandomBean {
+
+ private String value;
+ private Instant createdAt;
+ }
+}
diff --git a/src/test/java/dev/vality/testcontainers/annotations/util/SpringApplicationPropertiesLoaderTest.java b/src/test/java/dev/vality/testcontainers/annotations/util/SpringApplicationPropertiesLoaderTest.java
new file mode 100644
index 00000000..9e59c42a
--- /dev/null
+++ b/src/test/java/dev/vality/testcontainers/annotations/util/SpringApplicationPropertiesLoaderTest.java
@@ -0,0 +1,79 @@
+package dev.vality.testcontainers.annotations.util;
+
+import org.junit.jupiter.api.Test;
+import org.junit.jupiter.api.io.TempDir;
+
+import java.net.URLClassLoader;
+import java.nio.file.Files;
+import java.nio.file.Path;
+import java.util.List;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+class SpringApplicationPropertiesLoaderTest {
+
+ private static final String TOPIC_PROPERTY = "review.kafka.topic";
+
+ @TempDir
+ Path resources;
+
+ @Test
+ void shouldLoadActiveYamlDocumentAndProfileSpecificFile() throws Exception {
+ Files.writeString(resources.resolve("application.yml"), """
+ spring:
+ profiles:
+ active: review
+ review:
+ kafka:
+ topic: base-topic
+ ---
+ spring:
+ config:
+ activate:
+ on-profile: review
+ review:
+ kafka:
+ topic: profile-document-topic
+ """);
+ Files.writeString(resources.resolve("application-review.properties"),
+ TOPIC_PROPERTY + "=profile-file-topic\n");
+
+ var actual = withResourceClassLoader(() ->
+ SpringApplicationPropertiesLoader.loadFromSpringApplicationPropertiesFile(List.of(TOPIC_PROPERTY))
+ .getProperty(TOPIC_PROPERTY));
+
+ assertThat(actual).isEqualTo("profile-file-topic");
+ }
+
+ @Test
+ void shouldPreferSystemProperty() throws Exception {
+ Files.writeString(resources.resolve("application.properties"), TOPIC_PROPERTY + "=file-topic\n");
+ System.setProperty(TOPIC_PROPERTY, "system-topic");
+ try {
+ var actual = withResourceClassLoader(() ->
+ SpringApplicationPropertiesLoader.loadFromSpringApplicationPropertiesFile(List.of(TOPIC_PROPERTY))
+ .getProperty(TOPIC_PROPERTY));
+
+ assertThat(actual).isEqualTo("system-topic");
+ } finally {
+ System.clearProperty(TOPIC_PROPERTY);
+ }
+ }
+
+ private T withResourceClassLoader(ThrowingSupplier action) throws Exception {
+ var thread = Thread.currentThread();
+ var previous = thread.getContextClassLoader();
+ try (var classLoader = new URLClassLoader(new java.net.URL[] {resources.toUri().toURL()}, previous)) {
+ thread.setContextClassLoader(classLoader);
+ return action.get();
+ } finally {
+ thread.setContextClassLoader(previous);
+ }
+ }
+
+ @FunctionalInterface
+ private interface ThrowingSupplier {
+
+ T get() throws Exception;
+ }
+}
diff --git a/src/test/java/dev/vality/testcontainers/annotations/util/ValuesGeneratorTest.java b/src/test/java/dev/vality/testcontainers/annotations/util/ValuesGeneratorTest.java
new file mode 100644
index 00000000..92d8d41c
--- /dev/null
+++ b/src/test/java/dev/vality/testcontainers/annotations/util/ValuesGeneratorTest.java
@@ -0,0 +1,38 @@
+package dev.vality.testcontainers.annotations.util;
+
+import org.junit.jupiter.api.AfterEach;
+import org.junit.jupiter.api.Test;
+
+import java.time.Clock;
+import java.time.Instant;
+import java.time.ZoneId;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+class ValuesGeneratorTest {
+
+ @AfterEach
+ void resetClock() {
+ ValuesGenerator.resetClock();
+ }
+
+ @Test
+ void shouldAddCalendarDayAcrossDstTransition() {
+ var clock = Clock.fixed(
+ Instant.parse("2026-03-28T12:00:00Z"),
+ ZoneId.of("Europe/Amsterdam"));
+ ValuesGenerator.useClock(clock);
+
+ assertThat(ValuesGenerator.generateCurrentTimePlusDay())
+ .isEqualTo(Instant.parse("2026-03-29T11:00:00Z"));
+ }
+
+ @Test
+ void shouldCalculateTimeWindowAtCallTime() {
+ ValuesGenerator.useClock(Clock.fixed(Instant.parse("2026-07-25T10:00:00Z"), ZoneId.of("UTC")));
+
+ assertThat(ValuesGenerator.getFromTime().toString()).isEqualTo("2026-07-25T07:00");
+ assertThat(ValuesGenerator.getInFromToPeriodTime().toString()).isEqualTo("2026-07-25T08:00");
+ assertThat(ValuesGenerator.getToTime().toString()).isEqualTo("2026-07-25T09:00");
+ }
+}