diff --git a/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/api/tool/PaymentStatusFunctionCallingIT.java b/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/api/tool/PaymentStatusFunctionCallingIT.java index c542df6fc..f31d35c41 100644 --- a/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/api/tool/PaymentStatusFunctionCallingIT.java +++ b/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/api/tool/PaymentStatusFunctionCallingIT.java @@ -24,6 +24,7 @@ import java.util.function.Function; import com.fasterxml.jackson.annotation.JsonProperty; import com.fasterxml.jackson.core.JsonProcessingException; import com.fasterxml.jackson.databind.ObjectMapper; +import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.slf4j.Logger; @@ -52,6 +53,7 @@ import static org.assertj.core.api.Assertions.assertThat; * @author Christian Tzolov * @since 0.8.1 */ +@Disabled("See https://github.com/spring-projects/spring-ai/issues/1853") @EnabledIfEnvironmentVariable(named = "MISTRAL_AI_API_KEY", matches = ".+") public class PaymentStatusFunctionCallingIT { diff --git a/vector-stores/spring-ai-opensearch-store/src/test/java/org/springframework/ai/vectorstore/OpenSearchVectorStoreWithOllamaIT.java b/vector-stores/spring-ai-opensearch-store/src/test/java/org/springframework/ai/vectorstore/OpenSearchVectorStoreWithOllamaIT.java index bcccbeae1..ef71b5cf4 100644 --- a/vector-stores/spring-ai-opensearch-store/src/test/java/org/springframework/ai/vectorstore/OpenSearchVectorStoreWithOllamaIT.java +++ b/vector-stores/spring-ai-opensearch-store/src/test/java/org/springframework/ai/vectorstore/OpenSearchVectorStoreWithOllamaIT.java @@ -43,6 +43,9 @@ import org.springframework.ai.ollama.OllamaEmbeddingModel; import org.springframework.ai.ollama.api.OllamaApi; import org.springframework.ai.ollama.api.OllamaModel; import org.springframework.ai.ollama.api.OllamaOptions; +import org.springframework.ai.ollama.management.ModelManagementOptions; +import org.springframework.ai.ollama.management.OllamaModelManager; +import org.springframework.ai.ollama.management.PullModelStrategy; import org.springframework.beans.factory.annotation.Qualifier; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -56,6 +59,12 @@ import static org.hamcrest.Matchers.hasSize; @EnabledIfEnvironmentVariable(named = "OLLAMA_TESTS_ENABLED", matches = "true") class OpenSearchVectorStoreWithOllamaIT { + private static final Duration DEFAULT_TIMEOUT = Duration.ofMinutes(10); + + private static final int DEFAULT_MAX_RETRIES = 2; + + private static final String OLLAMA_LOCAL_URL = "http://localhost:11434"; + @Container private static final OpensearchContainer opensearchContainer = new OpensearchContainer<>( OpenSearchImage.DEFAULT_IMAGE); @@ -72,6 +81,19 @@ class OpenSearchVectorStoreWithOllamaIT { Awaitility.setDefaultPollInterval(2, TimeUnit.SECONDS); Awaitility.setDefaultPollDelay(Duration.ZERO); Awaitility.setDefaultTimeout(Duration.ofMinutes(1)); + + // Ensure the model is pulled before running tests + ensureModelIsPresent(OllamaModel.MXBAI_EMBED_LARGE.getName()); + } + + private static void ensureModelIsPresent(final String model) { + final OllamaApi api = new OllamaApi(OLLAMA_LOCAL_URL); + final var modelManagementOptions = ModelManagementOptions.builder() + .withMaxRetries(DEFAULT_MAX_RETRIES) + .withTimeout(DEFAULT_TIMEOUT) + .build(); + final var ollamaModelManager = new OllamaModelManager(api, modelManagementOptions); + ollamaModelManager.pullModel(model, PullModelStrategy.WHEN_MISSING); } private String getText(String uri) {