diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/BaseOllamaIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/BaseOllamaIT.java index 46abdadfd..53ebf3d27 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/BaseOllamaIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/BaseOllamaIT.java @@ -18,66 +18,87 @@ package org.springframework.ai.ollama; import java.time.Duration; +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.ollama.OllamaContainer; import org.springframework.ai.ollama.api.OllamaApi; import org.springframework.ai.ollama.management.ModelManagementOptions; import org.springframework.ai.ollama.management.OllamaModelManager; import org.springframework.ai.ollama.management.PullModelStrategy; -import org.springframework.util.StringUtils; +import org.springframework.util.Assert; -public class BaseOllamaIT { +@Testcontainers +@EnabledIfEnvironmentVariable(named = "OLLAMA_TESTS_ENABLED", matches = "true") +public abstract class BaseOllamaIT { - // Toggle for running tests locally on native Ollama for a faster feedback loop. - private static final boolean useTestcontainers = true; + private static final String OLLAMA_LOCAL_URL = "http://localhost:11434"; - public static OllamaContainer ollamaContainer; + private static final Duration DEFAULT_TIMEOUT = Duration.ofMinutes(10); - static { - if (useTestcontainers) { + private static final int DEFAULT_MAX_RETRIES = 2; + + // Environment variable to control whether to create a new container or use existing + // Ollama instance + private static final boolean SKIP_CONTAINER_CREATION = Boolean + .parseBoolean(System.getenv().getOrDefault("OLLAMA_SKIP_CONTAINER", "false")); + + private static OllamaContainer ollamaContainer; + + private static final ThreadLocal ollamaApi = new ThreadLocal<>(); + + /** + * Initialize the Ollama container and API with the specified model. This method + * should be called from @BeforeAll in subclasses. + * @param model the Ollama model to initialize (must not be null or empty) + * @return configured OllamaApi instance + * @throws IllegalArgumentException if model is null or empty + */ + protected static OllamaApi initializeOllama(final String model) { + Assert.hasText(model, "Model name must be provided"); + + if (!SKIP_CONTAINER_CREATION) { ollamaContainer = new OllamaContainer(OllamaImage.DEFAULT_IMAGE).withReuse(true); ollamaContainer.start(); } + + final OllamaApi api = buildOllamaApiWithModel(model); + ollamaApi.set(api); + return api; } /** - * Change the return value to false in order to run multiple Ollama IT tests locally - * reusing the same container image. - * - * Also, add the entry - * - * testcontainers.reuse.enable=true - * - * to the file ".testcontainers.properties" located in your home directory + * Get the initialized OllamaApi instance. + * @return the OllamaApi instance + * @throws IllegalStateException if called before initialization */ - public static boolean isDisabled() { - return true; + protected static OllamaApi getOllamaApi() { + OllamaApi api = ollamaApi.get(); + Assert.state(api != null, "OllamaApi not initialized. Call initializeOllama first."); + return api; } - public static OllamaApi buildOllamaApi() { - return buildOllamaApiWithModel(null); - } - - public static OllamaApi buildOllamaApiWithModel(String model) { - var baseUrl = "http://localhost:11434"; - if (useTestcontainers) { - baseUrl = ollamaContainer.getEndpoint(); + @AfterAll + public static void tearDown() { + if (ollamaContainer != null) { + ollamaContainer.stop(); } - var ollamaApi = new OllamaApi(baseUrl); - - if (StringUtils.hasText(model)) { - ensureModelIsPresent(ollamaApi, model); - } - - return ollamaApi; } - public static void ensureModelIsPresent(OllamaApi ollamaApi, String model) { - var modelManagementOptions = ModelManagementOptions.builder() - .withMaxRetries(2) - .withTimeout(Duration.ofMinutes(10)) + private static OllamaApi buildOllamaApiWithModel(final String model) { + final String baseUrl = SKIP_CONTAINER_CREATION ? OLLAMA_LOCAL_URL : ollamaContainer.getEndpoint(); + final OllamaApi api = new OllamaApi(baseUrl); + ensureModelIsPresent(api, model); + return api; + } + + private static void ensureModelIsPresent(final OllamaApi ollamaApi, final String model) { + final var modelManagementOptions = ModelManagementOptions.builder() + .withMaxRetries(DEFAULT_MAX_RETRIES) + .withTimeout(DEFAULT_TIMEOUT) .build(); - var ollamaModelManager = new OllamaModelManager(ollamaApi, modelManagementOptions); + final var ollamaModelManager = new OllamaModelManager(ollamaApi, modelManagementOptions); ollamaModelManager.pullModel(model, PullModelStrategy.WHEN_MISSING); } diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelFunctionCallingIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelFunctionCallingIT.java index c88e83db1..72d1b1b62 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelFunctionCallingIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelFunctionCallingIT.java @@ -22,10 +22,8 @@ import java.util.stream.Collectors; import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.testcontainers.junit.jupiter.Testcontainers; import reactor.core.publisher.Flux; import org.springframework.ai.chat.messages.AssistantMessage; @@ -46,9 +44,7 @@ import org.springframework.context.annotation.Bean; import static org.assertj.core.api.Assertions.assertThat; -@Testcontainers @SpringBootTest(classes = OllamaChatModelFunctionCallingIT.Config.class) -@DisabledIf("isDisabled") class OllamaChatModelFunctionCallingIT extends BaseOllamaIT { private static final Logger logger = LoggerFactory.getLogger(OllamaChatModelFunctionCallingIT.class); @@ -120,7 +116,7 @@ class OllamaChatModelFunctionCallingIT extends BaseOllamaIT { @Bean public OllamaApi ollamaApi() { - return buildOllamaApiWithModel(MODEL); + return initializeOllama(MODEL); } @Bean diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelIT.java index 88db3b96c..a5102586f 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelIT.java @@ -21,8 +21,6 @@ import java.util.Map; import java.util.stream.Collectors; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; -import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.chat.client.ChatClient; import org.springframework.ai.chat.messages.AssistantMessage; @@ -53,8 +51,6 @@ import org.springframework.core.convert.support.DefaultConversionService; import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest -@Testcontainers -@DisabledIf("isDisabled") class OllamaChatModelIT extends BaseOllamaIT { private static final String MODEL = OllamaModel.LLAMA3_2.getName(); @@ -241,7 +237,7 @@ class OllamaChatModelIT extends BaseOllamaIT { @Bean public OllamaApi ollamaApi() { - return buildOllamaApiWithModel(MODEL); + return initializeOllama(MODEL); } @Bean diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelMultimodalIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelMultimodalIT.java index 735857ef3..ee29deb67 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelMultimodalIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelMultimodalIT.java @@ -19,10 +19,8 @@ package org.springframework.ai.ollama; import java.util.List; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; @@ -40,8 +38,6 @@ import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.AssertionsForClassTypes.assertThatThrownBy; @SpringBootTest -@Testcontainers -@DisabledIf("isDisabled") class OllamaChatModelMultimodalIT extends BaseOllamaIT { private static final Logger logger = LoggerFactory.getLogger(OllamaChatModelMultimodalIT.class); @@ -80,7 +76,7 @@ class OllamaChatModelMultimodalIT extends BaseOllamaIT { @Bean public OllamaApi ollamaApi() { - return buildOllamaApiWithModel(MODEL); + return initializeOllama(MODEL); } @Bean diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelObservationIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelObservationIT.java index 5254f3e01..330ba7bcf 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelObservationIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelObservationIT.java @@ -23,7 +23,6 @@ import io.micrometer.observation.tck.TestObservationRegistry; import io.micrometer.observation.tck.TestObservationRegistryAssert; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; import reactor.core.publisher.Flux; import org.springframework.ai.chat.metadata.ChatResponseMetadata; @@ -50,7 +49,6 @@ import static org.springframework.ai.chat.observation.ChatModelObservationDocume * @author Thomas Vitale */ @SpringBootTest(classes = OllamaChatModelObservationIT.Config.class) -@DisabledIf("isDisabled") public class OllamaChatModelObservationIT extends BaseOllamaIT { private static final String MODEL = OllamaModel.LLAMA3_2.getName(); @@ -166,7 +164,7 @@ public class OllamaChatModelObservationIT extends BaseOllamaIT { @Bean public OllamaApi openAiApi() { - return buildOllamaApiWithModel(MODEL); + return initializeOllama(MODEL); } @Bean diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelIT.java index 00322204e..8853c1c85 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelIT.java @@ -19,8 +19,6 @@ package org.springframework.ai.ollama; import java.util.List; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; -import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.embedding.EmbeddingRequest; import org.springframework.ai.embedding.EmbeddingResponse; @@ -38,8 +36,6 @@ import org.springframework.context.annotation.Bean; import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest -@DisabledIf("isDisabled") -@Testcontainers class OllamaEmbeddingModelIT extends BaseOllamaIT { private static final String MODEL = OllamaModel.NOMIC_EMBED_TEXT.getName(); @@ -100,7 +96,7 @@ class OllamaEmbeddingModelIT extends BaseOllamaIT { @Bean public OllamaApi ollamaApi() { - return buildOllamaApiWithModel(MODEL); + return initializeOllama(MODEL); } @Bean diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelObservationIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelObservationIT.java index ad3ebc5a7..a25603cc0 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelObservationIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelObservationIT.java @@ -21,7 +21,6 @@ import java.util.List; import io.micrometer.observation.tck.TestObservationRegistry; import io.micrometer.observation.tck.TestObservationRegistryAssert; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; import org.springframework.ai.embedding.EmbeddingRequest; import org.springframework.ai.embedding.EmbeddingResponse; @@ -47,7 +46,6 @@ import static org.assertj.core.api.Assertions.assertThat; * @author Thomas Vitale */ @SpringBootTest(classes = OllamaEmbeddingModelObservationIT.Config.class) -@DisabledIf("isDisabled") public class OllamaEmbeddingModelObservationIT extends BaseOllamaIT { private static final String MODEL = OllamaModel.NOMIC_EMBED_TEXT.getName(); @@ -100,7 +98,7 @@ public class OllamaEmbeddingModelObservationIT extends BaseOllamaIT { @Bean public OllamaApi openAiApi() { - return buildOllamaApiWithModel(MODEL); + return initializeOllama(MODEL); } @Bean diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/OllamaApiIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/OllamaApiIT.java index cbf53fb3e..0b2b8cc7d 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/OllamaApiIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/OllamaApiIT.java @@ -22,8 +22,6 @@ import java.util.stream.Collectors; import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; -import org.testcontainers.junit.jupiter.Testcontainers; import reactor.core.publisher.Flux; import org.springframework.ai.ollama.BaseOllamaIT; @@ -42,17 +40,13 @@ import static org.assertj.core.api.Assertions.assertThat; * @author Christian Tzolov * @author Thomas Vitale */ -@Testcontainers -@DisabledIf("isDisabled") public class OllamaApiIT extends BaseOllamaIT { private static final String MODEL = OllamaModel.LLAMA3_2.getName(); - static OllamaApi ollamaApi; - @BeforeAll public static void beforeAll() throws IOException, InterruptedException { - ollamaApi = buildOllamaApiWithModel(MODEL); + initializeOllama(MODEL); } @Test @@ -63,7 +57,7 @@ public class OllamaApiIT extends BaseOllamaIT { .withStream(false) .build(); - GenerateResponse response = ollamaApi.generate(request); + GenerateResponse response = getOllamaApi().generate(request); System.out.println(response); @@ -87,7 +81,7 @@ public class OllamaApiIT extends BaseOllamaIT { .withOptions(OllamaOptions.create().withTemperature(0.9)) .build(); - ChatResponse response = ollamaApi.chat(request); + ChatResponse response = getOllamaApi().chat(request); System.out.println(response); @@ -108,7 +102,7 @@ public class OllamaApiIT extends BaseOllamaIT { .withOptions(OllamaOptions.create().withTemperature(0.9).toMap()) .build(); - Flux response = ollamaApi.streamingChat(request); + Flux response = getOllamaApi().streamingChat(request); List responses = response.collectList().block(); System.out.println(responses); @@ -128,7 +122,7 @@ public class OllamaApiIT extends BaseOllamaIT { public void embedText() { EmbeddingsRequest request = new EmbeddingsRequest(MODEL, "I like to eat apples"); - EmbeddingsResponse response = ollamaApi.embed(request); + EmbeddingsResponse response = getOllamaApi().embed(request); assertThat(response).isNotNull(); assertThat(response.embeddings()).hasSize(1); diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/OllamaApiModelsIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/OllamaApiModelsIT.java index f56525701..6a0886289 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/OllamaApiModelsIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/OllamaApiModelsIT.java @@ -21,8 +21,6 @@ import java.time.Duration; import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; -import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.ollama.BaseOllamaIT; import org.springframework.http.HttpStatus; @@ -34,8 +32,6 @@ import static org.assertj.core.api.Assertions.assertThat; * * @author Thomas Vitale */ -@Testcontainers -@DisabledIf("isDisabled") public class OllamaApiModelsIT extends BaseOllamaIT { private static final String MODEL = "all-minilm"; @@ -44,7 +40,7 @@ public class OllamaApiModelsIT extends BaseOllamaIT { @BeforeAll public static void beforeAll() throws IOException, InterruptedException { - ollamaApi = buildOllamaApiWithModel(MODEL); + ollamaApi = initializeOllama(MODEL); } @Test diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/tool/OllamaApiToolFunctionCallIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/tool/OllamaApiToolFunctionCallIT.java index 81d1f56b5..7ca5b57e3 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/tool/OllamaApiToolFunctionCallIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/api/tool/OllamaApiToolFunctionCallIT.java @@ -23,10 +23,8 @@ import java.util.Map; import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.ollama.BaseOllamaIT; @@ -42,8 +40,6 @@ import static org.assertj.core.api.Assertions.assertThat; * @author Christian Tzolov * @author Thomas Vitale */ -@Testcontainers -@DisabledIf("isDisabled") public class OllamaApiToolFunctionCallIT extends BaseOllamaIT { private static final String MODEL = "qwen2.5:3b"; @@ -56,7 +52,7 @@ public class OllamaApiToolFunctionCallIT extends BaseOllamaIT { @BeforeAll public static void beforeAll() throws IOException, InterruptedException { - ollamaApi = buildOllamaApiWithModel(MODEL); + ollamaApi = initializeOllama(MODEL); } @SuppressWarnings("null") diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/management/OllamaModelManagerIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/management/OllamaModelManagerIT.java index 2d640585b..eeb6e7c62 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/management/OllamaModelManagerIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/management/OllamaModelManagerIT.java @@ -22,8 +22,6 @@ import java.util.List; import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; -import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.ollama.BaseOllamaIT; import org.springframework.ai.ollama.api.OllamaModel; @@ -35,8 +33,6 @@ import static org.assertj.core.api.Assertions.assertThat; * * @author Thomas Vitale */ -@Testcontainers -@DisabledIf("isDisabled") class OllamaModelManagerIT extends BaseOllamaIT { private static final String MODEL = OllamaModel.NOMIC_EMBED_TEXT.getName(); @@ -45,7 +41,7 @@ class OllamaModelManagerIT extends BaseOllamaIT { @BeforeAll public static void beforeAll() throws IOException, InterruptedException { - var ollamaApi = buildOllamaApiWithModel(MODEL); + var ollamaApi = initializeOllama(MODEL); modelManager = new OllamaModelManager(ollamaApi); } @@ -144,7 +140,7 @@ class OllamaModelManagerIT extends BaseOllamaIT { var isModelAvailable = modelManager.isModelAvailable(model); assertThat(isModelAvailable).isFalse(); - new OllamaModelManager(buildOllamaApi(), + new OllamaModelManager(getOllamaApi(), new ModelManagementOptions(PullModelStrategy.WHEN_MISSING, List.of(model), Duration.ofMinutes(5), 0)); isModelAvailable = modelManager.isModelAvailable(model); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/BaseOllamaIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/BaseOllamaIT.java index e27cc47e4..b4dbc0ee9 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/BaseOllamaIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/BaseOllamaIT.java @@ -19,27 +19,65 @@ package org.springframework.ai.autoconfigure.ollama; import java.time.Duration; import org.junit.jupiter.api.AfterAll; -import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.ollama.OllamaContainer; import org.springframework.ai.ollama.api.OllamaApi; import org.springframework.ai.ollama.management.ModelManagementOptions; import org.springframework.ai.ollama.management.OllamaModelManager; import org.springframework.ai.ollama.management.PullModelStrategy; +import org.springframework.util.Assert; -public class BaseOllamaIT { +@Testcontainers +@EnabledIfEnvironmentVariable(named = "OLLAMA_TESTS_ENABLED", matches = "true") +public abstract class BaseOllamaIT { + + private static final String OLLAMA_LOCAL_URL = "http://localhost:11434"; + + private static final Duration DEFAULT_TIMEOUT = Duration.ofMinutes(10); + + private static final int DEFAULT_MAX_RETRIES = 2; + + // Environment variable to control whether to create a new container or use existing + // Ollama instance + private static final boolean SKIP_CONTAINER_CREATION = Boolean + .parseBoolean(System.getenv().getOrDefault("OLLAMA_WITH_REUSE", "false")); private static OllamaContainer ollamaContainer; - // Toggle for running tests locally on native Ollama for a faster feedback loop. - private static final boolean useTestcontainers = true; + private static final ThreadLocal ollamaApi = new ThreadLocal<>(); - @BeforeAll - public static void setUp() { - if (useTestcontainers && !isDisabled()) { - ollamaContainer = new OllamaContainer(OllamaImage.IMAGE).withReuse(true); + /** + * Initialize the Ollama API with the specified model. When OLLAMA_WITH_REUSE=true + * (default), uses TestContainers withReuse feature. When OLLAMA_WITH_REUSE=false, + * connects to local Ollama instance. + * @param model the Ollama model to initialize (must not be null or empty) + * @return configured OllamaApi instance + * @throws IllegalArgumentException if model is null or empty + */ + protected static OllamaApi initializeOllama(final String model) { + Assert.hasText(model, "Model name must be provided"); + + if (!SKIP_CONTAINER_CREATION) { + ollamaContainer = new OllamaContainer(OllamaImage.DEFAULT_IMAGE).withReuse(true); ollamaContainer.start(); } + + final OllamaApi api = buildOllamaApiWithModel(model); + ollamaApi.set(api); + return api; + } + + /** + * Get the initialized OllamaApi instance. + * @return the OllamaApi instance + * @throws IllegalStateException if called before initialization + */ + protected static OllamaApi getOllamaApi() { + OllamaApi api = ollamaApi.get(); + Assert.state(api != null, "OllamaApi not initialized. Call initializeOllama first."); + return api; } @AfterAll @@ -49,33 +87,25 @@ public class BaseOllamaIT { } } - /** - * Change the return value to false in order to run multiple Ollama IT tests locally - * reusing the same container image. - * - * Also, add the entry - * - * testcontainers.reuse.enable=true - * - * to the file ".testcontainers.properties" located in your home directory - */ - public static boolean isDisabled() { - return true; + public static OllamaApi buildOllamaApiWithModel(final String model) { + final String baseUrl = SKIP_CONTAINER_CREATION ? OLLAMA_LOCAL_URL : ollamaContainer.getEndpoint(); + final OllamaApi api = new OllamaApi(baseUrl); + ensureModelIsPresent(api, model); + return api; } - public static String buildConnectionWithModel(String model) { - var baseUrl = "http://localhost:11434"; - if (useTestcontainers) { - baseUrl = ollamaContainer.getEndpoint(); - } - - var modelManagementOptions = ModelManagementOptions.builder() - .withMaxRetries(2) - .withTimeout(Duration.ofMinutes(10)) - .build(); - var ollamaModelManager = new OllamaModelManager(new OllamaApi(baseUrl), modelManagementOptions); - ollamaModelManager.pullModel(model, PullModelStrategy.WHEN_MISSING); + public String getBaseUrl() { + String baseUrl = SKIP_CONTAINER_CREATION ? OLLAMA_LOCAL_URL : ollamaContainer.getEndpoint(); return baseUrl; } + private static void ensureModelIsPresent(final OllamaApi ollamaApi, final String model) { + final var modelManagementOptions = ModelManagementOptions.builder() + .withMaxRetries(DEFAULT_MAX_RETRIES) + .withTimeout(DEFAULT_TIMEOUT) + .build(); + final var ollamaModelManager = new OllamaModelManager(ollamaApi, modelManagementOptions); + ollamaModelManager.pullModel(model, PullModelStrategy.WHEN_MISSING); + } + } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaChatAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaChatAutoConfigurationIT.java index b596b14f0..335378099 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaChatAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaChatAutoConfigurationIT.java @@ -22,8 +22,6 @@ import java.util.stream.Collectors; import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; -import org.testcontainers.junit.jupiter.Testcontainers; import reactor.core.publisher.Flux; import org.springframework.ai.chat.messages.AssistantMessage; @@ -46,17 +44,13 @@ import static org.assertj.core.api.Assertions.assertThat; * @author Thomas Vitale * @since 0.8.0 */ -@Testcontainers -@DisabledIf("isDisabled") public class OllamaChatAutoConfigurationIT extends BaseOllamaIT { private static final String MODEL_NAME = OllamaModel.LLAMA3_2.getName(); - static String baseUrl; - private final ApplicationContextRunner contextRunner = new ApplicationContextRunner().withPropertyValues( // @formatter:off - "spring.ai.ollama.baseUrl=" + baseUrl, + "spring.ai.ollama.baseUrl=" + getBaseUrl(), "spring.ai.ollama.chat.options.model=" + MODEL_NAME, "spring.ai.ollama.chat.options.temperature=0.5", "spring.ai.ollama.chat.options.topK=10") @@ -67,7 +61,7 @@ public class OllamaChatAutoConfigurationIT extends BaseOllamaIT { @BeforeAll public static void beforeAll() throws IOException, InterruptedException { - baseUrl = buildConnectionWithModel(MODEL_NAME); + initializeOllama(MODEL_NAME); } @Test diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaEmbeddingAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaEmbeddingAutoConfigurationIT.java index 0a2521aef..fc6fd6770 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaEmbeddingAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaEmbeddingAutoConfigurationIT.java @@ -21,8 +21,6 @@ import java.util.List; import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; -import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.embedding.EmbeddingResponse; import org.springframework.ai.ollama.OllamaEmbeddingModel; @@ -38,24 +36,20 @@ import static org.assertj.core.api.Assertions.assertThat; /** * @author Christian Tzolov * @author Thomas Vitale - * @since 0.8.0 + * @since 1.0.0 */ -@Testcontainers -@DisabledIf("isDisabled") public class OllamaEmbeddingAutoConfigurationIT extends BaseOllamaIT { private static final String MODEL_NAME = OllamaModel.NOMIC_EMBED_TEXT.getName(); - static String baseUrl; - private final ApplicationContextRunner contextRunner = new ApplicationContextRunner() .withPropertyValues("spring.ai.ollama.embedding.options.model=" + MODEL_NAME, - "spring.ai.ollama.base-url=" + baseUrl) + "spring.ai.ollama.base-url=" + getBaseUrl()) .withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OllamaAutoConfiguration.class)); @BeforeAll public static void beforeAll() throws IOException, InterruptedException { - baseUrl = buildConnectionWithModel(MODEL_NAME); + initializeOllama(MODEL_NAME); } @Test diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaImage.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaImage.java index 75d243ecd..c95c80fab 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaImage.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaImage.java @@ -18,7 +18,7 @@ package org.springframework.ai.autoconfigure.ollama; public final class OllamaImage { - public static final String IMAGE = "ollama/ollama:0.3.14"; + public static final String DEFAULT_IMAGE = "ollama/ollama:0.3.14"; private OllamaImage() { diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackInPromptIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackInPromptIT.java index 57c388962..047df41f7 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackInPromptIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackInPromptIT.java @@ -22,10 +22,8 @@ import java.util.stream.Collectors; import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.testcontainers.junit.jupiter.Testcontainers; import reactor.core.publisher.Flux; import org.springframework.ai.autoconfigure.ollama.BaseOllamaIT; @@ -43,19 +41,15 @@ import org.springframework.boot.test.context.runner.ApplicationContextRunner; import static org.assertj.core.api.Assertions.assertThat; -@Testcontainers -@DisabledIf("isDisabled") public class FunctionCallbackInPromptIT extends BaseOllamaIT { private static final Logger logger = LoggerFactory.getLogger(FunctionCallbackInPromptIT.class); private static final String MODEL_NAME = "qwen2.5:3b"; - static String baseUrl; - private final ApplicationContextRunner contextRunner = new ApplicationContextRunner().withPropertyValues( // @formatter:off - "spring.ai.ollama.baseUrl=" + baseUrl, + "spring.ai.ollama.baseUrl=" + getBaseUrl(), "spring.ai.ollama.chat.options.model=" + MODEL_NAME, "spring.ai.ollama.chat.options.temperature=0.5", "spring.ai.ollama.chat.options.topK=10") @@ -64,7 +58,7 @@ public class FunctionCallbackInPromptIT extends BaseOllamaIT { @BeforeAll public static void beforeAll() { - baseUrl = buildConnectionWithModel(MODEL_NAME); + initializeOllama(MODEL_NAME); } @Test diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackWrapperIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackWrapperIT.java index ce686de12..d500512c0 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackWrapperIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackWrapperIT.java @@ -22,10 +22,8 @@ import java.util.stream.Collectors; import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.DisabledIf; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.testcontainers.junit.jupiter.Testcontainers; import reactor.core.publisher.Flux; import org.springframework.ai.autoconfigure.ollama.BaseOllamaIT; @@ -48,19 +46,15 @@ import org.springframework.context.annotation.Configuration; import static org.assertj.core.api.Assertions.assertThat; -@Testcontainers -@DisabledIf("isDisabled") public class FunctionCallbackWrapperIT extends BaseOllamaIT { private static final Logger logger = LoggerFactory.getLogger(FunctionCallbackWrapperIT.class); private static final String MODEL_NAME = "qwen2.5:3b"; - static String baseUrl; - private final ApplicationContextRunner contextRunner = new ApplicationContextRunner().withPropertyValues( // @formatter:off - "spring.ai.ollama.baseUrl=" + baseUrl, + "spring.ai.ollama.baseUrl=" + getBaseUrl(), "spring.ai.ollama.chat.options.model=" + MODEL_NAME, "spring.ai.ollama.chat.options.temperature=0.5", "spring.ai.ollama.chat.options.topK=10") @@ -70,7 +64,7 @@ public class FunctionCallbackWrapperIT extends BaseOllamaIT { @BeforeAll public static void beforeAll() { - baseUrl = buildConnectionWithModel(MODEL_NAME); + initializeOllama(MODEL_NAME); } @Test diff --git a/spring-ai-spring-boot-autoconfigure/src/test/kotlin/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackContextKotlinIT.kt b/spring-ai-spring-boot-autoconfigure/src/test/kotlin/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackContextKotlinIT.kt index 8866b2948..3e1cb90aa 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/kotlin/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackContextKotlinIT.kt +++ b/spring-ai-spring-boot-autoconfigure/src/test/kotlin/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackContextKotlinIT.kt @@ -43,7 +43,7 @@ class FunctionCallbackContextKotlinIT : BaseOllamaIT() { private val MODEL_NAME = "qwen2.5:3b" - val contextRunner = buildConnectionWithModel(MODEL_NAME).let { baseUrl -> + val contextRunner = buildOllamaApiWithModel(MODEL_NAME).let { baseUrl -> ApplicationContextRunner().withPropertyValues( "spring.ai.ollama.baseUrl=$baseUrl", "spring.ai.ollama.chat.options.model=$MODEL_NAME", diff --git a/spring-ai-spring-boot-autoconfigure/src/test/kotlin/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackWrapperKotlinIT.kt b/spring-ai-spring-boot-autoconfigure/src/test/kotlin/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackWrapperKotlinIT.kt index 4af4d38c7..fca6c96be 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/kotlin/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackWrapperKotlinIT.kt +++ b/spring-ai-spring-boot-autoconfigure/src/test/kotlin/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackWrapperKotlinIT.kt @@ -44,7 +44,7 @@ class FunctionCallbackWrapperKotlinIT : BaseOllamaIT() { private val MODEL_NAME = "qwen2.5:3b" - val contextRunner = buildConnectionWithModel(MODEL_NAME).let { baseUrl -> + val contextRunner = buildOllamaApiWithModel(MODEL_NAME).let { baseUrl -> ApplicationContextRunner().withPropertyValues( "spring.ai.ollama.baseUrl=$baseUrl", "spring.ai.ollama.chat.options.model=$MODEL_NAME",