diff --git a/modules/database-commons/src/main/java/org/testcontainers/ext/ScriptUtils.java b/modules/database-commons/src/main/java/org/testcontainers/ext/ScriptUtils.java index b78ac4d63c4..15f30d6fa34 100644 --- a/modules/database-commons/src/main/java/org/testcontainers/ext/ScriptUtils.java +++ b/modules/database-commons/src/main/java/org/testcontainers/ext/ScriptUtils.java @@ -28,6 +28,8 @@ import java.util.LinkedList; import java.util.List; import java.util.concurrent.TimeUnit; +import java.util.function.UnaryOperator; +import java.util.regex.Pattern; import javax.script.ScriptException; @@ -69,6 +71,15 @@ public abstract class ScriptUtils { */ public static final String DEFAULT_BLOCK_COMMENT_END_DELIMITER = "*/"; + /** + * T-SQL batch separator used by Microsoft SQL Server tooling such as {@code sqlcmd} and SSMS. + * Pass this as the {@code separator} argument when executing scripts that use {@code GO} as a batch + * delimiter instead of {@code ;}. + */ + public static final String GO_STATEMENT_SEPARATOR = "GO"; + + private static final Pattern GO_SEPARATOR_PATTERN = Pattern.compile("(?im)^[ \\t]*GO[ \\t]*$"); + /** * Prevent instantiation of this utility class. */ @@ -191,6 +202,52 @@ public static boolean containsSqlScriptDelimiters( return false; } + /** + * Replaces standalone {@code GO} batch separators with {@code ;} so that T-SQL scripts produced + * by tools such as {@code sqlcmd} or SSMS can be fed to the standard script executor. + * + * @param script the raw SQL script content + * @return the script with {@code GO} separators replaced by {@code ;} + */ + public static String normalizeGoSeparator(String script) { + return GO_SEPARATOR_PATTERN.matcher(script).replaceAll(";"); + } + + /** + * Load script from classpath, apply a preprocessor, and execute it against the given database. + * + * @param databaseDelegate database delegate for script execution + * @param initScriptPath the resource to load the init script from + * @param scriptPreprocessor function applied to the raw script content before execution + */ + public static void runInitScript( + DatabaseDelegate databaseDelegate, + String initScriptPath, + UnaryOperator scriptPreprocessor + ) { + try { + URL resource = Thread.currentThread().getContextClassLoader().getResource(initScriptPath); + if (resource == null) { + resource = ScriptUtils.class.getClassLoader().getResource(initScriptPath); + if (resource == null) { + LOGGER.warn("Could not load classpath init script: {}", initScriptPath); + throw new ScriptLoadException( + "Could not load classpath init script: " + initScriptPath + ". Resource not found." + ); + } + } + String scripts = IOUtils.toString(resource, StandardCharsets.UTF_8); + scripts = scriptPreprocessor.apply(scripts); + executeDatabaseScript(databaseDelegate, initScriptPath, scripts); + } catch (IOException e) { + LOGGER.warn("Could not load classpath init script: {}", initScriptPath); + throw new ScriptLoadException("Could not load classpath init script: " + initScriptPath, e); + } catch (ScriptException e) { + LOGGER.error("Error while executing init script: {}", initScriptPath, e); + throw new UncategorizedScriptException("Error while executing init script: " + initScriptPath, e); + } + } + /** * Load script from classpath and apply it to the given database * diff --git a/modules/database-commons/src/test/java/org/testcontainers/ext/ScriptSplittingTest.java b/modules/database-commons/src/test/java/org/testcontainers/ext/ScriptSplittingTest.java index cb0e33162f4..a8a0cf508fc 100644 --- a/modules/database-commons/src/test/java/org/testcontainers/ext/ScriptSplittingTest.java +++ b/modules/database-commons/src/test/java/org/testcontainers/ext/ScriptSplittingTest.java @@ -478,4 +478,36 @@ void testIgnoreDelimitersInLiteralsAndComments() { void testContainsDelimiters() { assertThat(ScriptUtils.containsSqlScriptDelimiters("'@' /*@*/ @ \"@\" --@", "@")).isTrue(); } + + @Test + void testNormalizeGoSeparatorBasic() { + String script = "SELECT 1\nGO\nSELECT 2\nGO\n"; + String normalized = ScriptUtils.normalizeGoSeparator(script); + List statements = doSplit(normalized, ScriptUtils.DEFAULT_STATEMENT_SEPARATOR); + assertThat(statements).containsExactly("SELECT 1", "SELECT 2"); + } + + @Test + void testNormalizeGoSeparatorCaseInsensitive() { + String script = "SELECT 1\ngo\nSELECT 2\nGo\n"; + String normalized = ScriptUtils.normalizeGoSeparator(script); + List statements = doSplit(normalized, ScriptUtils.DEFAULT_STATEMENT_SEPARATOR); + assertThat(statements).containsExactly("SELECT 1", "SELECT 2"); + } + + @Test + void testNormalizeGoSeparatorWithLeadingWhitespace() { + String script = "SELECT 1\n GO\nSELECT 2\n\tGO\n"; + String normalized = ScriptUtils.normalizeGoSeparator(script); + List statements = doSplit(normalized, ScriptUtils.DEFAULT_STATEMENT_SEPARATOR); + assertThat(statements).containsExactly("SELECT 1", "SELECT 2"); + } + + @Test + void testNormalizeGoSeparatorDoesNotMatchInlineGo() { + String script = "SELECT GOOD, GOTO_COL\nGO\n"; + String normalized = ScriptUtils.normalizeGoSeparator(script); + List statements = doSplit(normalized, ScriptUtils.DEFAULT_STATEMENT_SEPARATOR); + assertThat(statements).containsExactly("SELECT GOOD, GOTO_COL"); + } } diff --git a/modules/jdbc/src/main/java/org/testcontainers/containers/JdbcDatabaseContainer.java b/modules/jdbc/src/main/java/org/testcontainers/containers/JdbcDatabaseContainer.java index cf6c995528f..e307f9d68d3 100644 --- a/modules/jdbc/src/main/java/org/testcontainers/containers/JdbcDatabaseContainer.java +++ b/modules/jdbc/src/main/java/org/testcontainers/containers/JdbcDatabaseContainer.java @@ -358,6 +358,17 @@ protected void optionallyMapResourceParameterAsVolume( } } + /** + * Override to preprocess the raw SQL script content before execution. + * The default implementation returns the script unchanged. + * + * @param script the raw script content + * @return the preprocessed script content + */ + protected String preprocessInitScript(String script) { + return script; + } + /** * Load init script content and apply it to the database if initScriptPath is set */ @@ -365,7 +376,7 @@ protected void runInitScriptIfRequired() { initScriptPaths .stream() .filter(Objects::nonNull) - .forEach(path -> ScriptUtils.runInitScript(getDatabaseDelegate(), path)); + .forEach(path -> ScriptUtils.runInitScript(getDatabaseDelegate(), path, this::preprocessInitScript)); } public void setParameters(Map parameters) { diff --git a/modules/mssqlserver/src/main/java/org/testcontainers/containers/MSSQLServerContainer.java b/modules/mssqlserver/src/main/java/org/testcontainers/containers/MSSQLServerContainer.java index 07f8a064c00..8cb4dc890e0 100644 --- a/modules/mssqlserver/src/main/java/org/testcontainers/containers/MSSQLServerContainer.java +++ b/modules/mssqlserver/src/main/java/org/testcontainers/containers/MSSQLServerContainer.java @@ -1,5 +1,6 @@ package org.testcontainers.containers; +import org.testcontainers.ext.ScriptUtils; import org.testcontainers.utility.DockerImageName; import org.testcontainers.utility.LicenseAcceptance; @@ -166,4 +167,9 @@ private void checkPasswordStrength(String password) { ); } } + + @Override + protected String preprocessInitScript(String script) { + return ScriptUtils.normalizeGoSeparator(script); + } } diff --git a/modules/mssqlserver/src/main/java/org/testcontainers/mssqlserver/MSSQLServerContainer.java b/modules/mssqlserver/src/main/java/org/testcontainers/mssqlserver/MSSQLServerContainer.java index 6ad74ea7e72..5afdf07eca3 100644 --- a/modules/mssqlserver/src/main/java/org/testcontainers/mssqlserver/MSSQLServerContainer.java +++ b/modules/mssqlserver/src/main/java/org/testcontainers/mssqlserver/MSSQLServerContainer.java @@ -1,6 +1,7 @@ package org.testcontainers.mssqlserver; import org.testcontainers.containers.JdbcDatabaseContainer; +import org.testcontainers.ext.ScriptUtils; import org.testcontainers.utility.DockerImageName; import org.testcontainers.utility.LicenseAcceptance; @@ -153,4 +154,9 @@ private void checkPasswordStrength(String password) { ); } } + + @Override + protected String preprocessInitScript(String script) { + return ScriptUtils.normalizeGoSeparator(script); + } }