diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java index baeeecd3f..f2f5a168c 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java @@ -22,7 +22,6 @@ import java.util.Map; import java.util.Set; import org.springframework.ai.chat.messages.AssistantMessage; -import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.chat.messages.SystemMessage; import org.springframework.ai.chat.messages.ToolResponseMessage; import org.springframework.ai.chat.messages.UserMessage; @@ -59,6 +58,7 @@ import reactor.core.publisher.Flux; * * @author Christian Tzolov * @author luocongqiu + * @author Thomas Vitale * @since 1.0.0 */ public class OllamaChatModel extends AbstractToolCallSupport implements ChatModel { @@ -125,13 +125,13 @@ public class OllamaChatModel extends AbstractToolCallSupport implements ChatMode ChatGenerationMetadata generationMetadata = ChatGenerationMetadata.NULL; if (response.promptEvalCount() != null && response.evalCount() != null) { - generationMetadata = ChatGenerationMetadata.from("DONE", null); + generationMetadata = ChatGenerationMetadata.from(response.doneReason(), null); } var generator = new Generation(assistantMessage, generationMetadata); var chatResponse = new ChatResponse(List.of(generator), from(response)); - if (isToolCall(chatResponse, Set.of("DONE"))) { + if (isToolCall(chatResponse, Set.of("stop"))) { var toolCallConversation = handleToolCalls(prompt, chatResponse); // Recursively call the call method with the tool call message // conversation that contains the call responses. @@ -176,7 +176,7 @@ public class OllamaChatModel extends AbstractToolCallSupport implements ChatMode ChatGenerationMetadata generationMetadata = ChatGenerationMetadata.NULL; if (chunk.promptEvalCount() != null && chunk.evalCount() != null) { - generationMetadata = ChatGenerationMetadata.from("DONE", null); + generationMetadata = ChatGenerationMetadata.from(chunk.doneReason(), null); } var generator = new Generation(assistantMessage, generationMetadata); @@ -184,7 +184,7 @@ public class OllamaChatModel extends AbstractToolCallSupport implements ChatMode }); return chatResponse.flatMap(response -> { - if (isToolCall(response, Set.of("DONE"))) { + if (isToolCall(response, Set.of("stop"))) { var toolCallConversation = handleToolCalls(prompt, response); // Recursively call the stream method with the tool call message // conversation that contains the call responses. @@ -201,53 +201,43 @@ public class OllamaChatModel extends AbstractToolCallSupport implements ChatMode */ OllamaApi.ChatRequest ollamaChatRequest(Prompt prompt, boolean stream) { - List ollamaMessages = prompt.getInstructions() - .stream() - .filter(message -> message.getMessageType() == MessageType.USER - || message.getMessageType() == MessageType.ASSISTANT - || message.getMessageType() == MessageType.SYSTEM || message.getMessageType() == MessageType.TOOL) - .map(message -> { - if (message instanceof UserMessage userMessage) { - var messageBuilder = OllamaApi.Message.builder(Role.USER).withContent(message.getContent()); - if (!CollectionUtils.isEmpty(userMessage.getMedia())) { - messageBuilder.withImages(userMessage.getMedia() - .stream() - .map(media -> this.fromMediaData(media.getData())) - .toList()); - } - return List.of(messageBuilder.build()); + List ollamaMessages = prompt.getInstructions().stream().map(message -> { + if (message instanceof UserMessage userMessage) { + var messageBuilder = OllamaApi.Message.builder(Role.USER).withContent(message.getContent()); + if (!CollectionUtils.isEmpty(userMessage.getMedia())) { + messageBuilder.withImages( + userMessage.getMedia().stream().map(media -> this.fromMediaData(media.getData())).toList()); } - else if (message instanceof SystemMessage systemMessage) { - return List - .of(OllamaApi.Message.builder(Role.SYSTEM).withContent(systemMessage.getContent()).build()); + return List.of(messageBuilder.build()); + } + else if (message instanceof SystemMessage systemMessage) { + return List.of(OllamaApi.Message.builder(Role.SYSTEM).withContent(systemMessage.getContent()).build()); + } + else if (message instanceof AssistantMessage assistantMessage) { + List toolCalls = null; + if (!CollectionUtils.isEmpty(assistantMessage.getToolCalls())) { + toolCalls = assistantMessage.getToolCalls().stream().map(toolCall -> { + var function = new ToolCallFunction(toolCall.name(), + ModelOptionsUtils.jsonToMap(toolCall.arguments())); + return new ToolCall(function); + }).toList(); } - else if (message instanceof AssistantMessage assistantMessage) { - List toolCalls = null; - if (!CollectionUtils.isEmpty(assistantMessage.getToolCalls())) { - toolCalls = assistantMessage.getToolCalls().stream().map(toolCall -> { - var function = new ToolCallFunction(toolCall.name(), - ModelOptionsUtils.jsonToMap(toolCall.arguments())); - return new ToolCall(function); - }).toList(); - } - return List.of(OllamaApi.Message.builder(Role.ASSISTANT) - .withContent(assistantMessage.getContent()) - .withToolCalls(toolCalls) - .build()); - } - else if (message instanceof ToolResponseMessage toolMessage) { + return List.of(OllamaApi.Message.builder(Role.ASSISTANT) + .withContent(assistantMessage.getContent()) + .withToolCalls(toolCalls) + .build()); + } + else if (message instanceof ToolResponseMessage toolMessage) { - List responseMessages = toolMessage.getResponses() - .stream() - .map(tr -> OllamaApi.Message.builder(Role.TOOL).withContent(tr.responseData()).build()) - .toList(); + List responseMessages = toolMessage.getResponses() + .stream() + .map(tr -> OllamaApi.Message.builder(Role.TOOL).withContent(tr.responseData()).build()) + .toList(); - return responseMessages; - } - throw new IllegalArgumentException("Unsupported message type: " + message.getMessageType()); - }) - .flatMap(List::stream) - .toList(); + return responseMessages; + } + throw new IllegalArgumentException("Unsupported message type: " + message.getMessageType()); + }).flatMap(List::stream).toList(); Set functionsForThisRequest = new HashSet<>(); diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaApi.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaApi.java index 83acf37a7..a82f60817 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaApi.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaApi.java @@ -47,6 +47,7 @@ import reactor.core.publisher.Mono; * Java Client for the Ollama API. https://ollama.ai * * @author Christian Tzolov + * @author Thomas Vitale * @since 0.8.0 */ // @formatter:off @@ -454,15 +455,20 @@ public class OllamaApi { /** * Chat request object. * - * @param model The model to use for completion. - * @param messages The list of messages to chat with. - * @param stream Whether to stream the response. - * @param format The format to return the response in. Currently, the only accepted - * value is "json". - * @param keepAlive The duration to keep the model loaded in ollama while idle. - * @param options Additional model parameters. You can use the {@link OllamaOptions} builder - * to create the options then {@link OllamaOptions#toMap()} to convert the options into a - * map. + * @param model The model to use for completion. It should be a name familiar to Ollama from the Library. + * @param messages The list of messages in the chat. This can be used to keep a chat memory. + * @param stream Whether to stream the response. If false, the response will be returned as a single response object rather than a stream of objects. + * @param format The format to return the response in. Currently, the only accepted value is "json". + * @param keepAlive Controls how long the model will stay loaded into memory following this request (default: 5m). + * @param tools List of tools the model has access to. + * @param options Model-specific options. For example, "temperature" can be set through this field, if the model supports it. + * You can use the {@link OllamaOptions} builder to create the options then {@link OllamaOptions#toMap()} to convert the options into a map. + * + * @see Chat + * Completion API + * @see Ollama + * Types */ @JsonInclude(Include.NON_NULL) public record ChatRequest( @@ -471,9 +477,9 @@ public class OllamaApi { @JsonProperty("stream") Boolean stream, @JsonProperty("format") String format, @JsonProperty("keep_alive") String keepAlive, - @JsonProperty("options") Map options, - @JsonProperty("tools") List tools) { - + @JsonProperty("tools") List tools, + @JsonProperty("options") Map options + ) { /** * Represents a tool the model may call. Currently, only functions are supported as a tool. @@ -544,8 +550,8 @@ public class OllamaApi { private boolean stream = false; private String format; private String keepAlive; - private Map options = Map.of(); private List tools = List.of(); + private Map options = Map.of(); public Builder(String model) { Assert.notNull(model, "The model can not be null."); @@ -572,6 +578,11 @@ public class OllamaApi { return this; } + public Builder withTools(List tools) { + this.tools = tools; + return this; + } + public Builder withOptions(Map options) { Objects.requireNonNull(options, "The options can not be null."); @@ -585,13 +596,8 @@ public class OllamaApi { return this; } - public Builder withTools(List tools) { - this.tools = tools; - return this; - } - public ChatRequest build() { - return new ChatRequest(model, messages, stream, format, keepAlive, options, tools); + return new ChatRequest(model, messages, stream, format, keepAlive, tools, options); } } } @@ -599,19 +605,21 @@ public class OllamaApi { /** * Ollama chat response object. * - * @param model The model name used for completion. - * @param createdAt When the request was made. + * @param model The model used for generating the response. + * @param createdAt The timestamp of the response generation. * @param message The response {@link Message} with {@link Message.Role#ASSISTANT}. + * @param doneReason The reason the model stopped generating text. * @param done Whether this is the final response. For streaming response only the * last message is marked as done. If true, this response may be followed by another * response with the following, additional fields: context, prompt_eval_count, * prompt_eval_duration, eval_count, eval_duration. * @param totalDuration Time spent generating the response. * @param loadDuration Time spent loading the model. - * @param promptEvalCount number of tokens in the prompt.(*) - * @param promptEvalDuration time spent evaluating the prompt. - * @param evalCount number of tokens in the response. - * @param evalDuration time spent generating the response. + * @param promptEvalCount Number of tokens in the prompt. + * @param promptEvalDuration Time spent evaluating the prompt. + * @param evalCount Number of tokens in the response. + * @param evalDuration Time spent generating the response. + * * @see Chat * Completion API @@ -623,13 +631,15 @@ public class OllamaApi { @JsonProperty("model") String model, @JsonProperty("created_at") Instant createdAt, @JsonProperty("message") Message message, + @JsonProperty("done_reason") String doneReason, @JsonProperty("done") Boolean done, @JsonProperty("total_duration") Duration totalDuration, @JsonProperty("load_duration") Duration loadDuration, @JsonProperty("prompt_eval_count") Integer promptEvalCount, @JsonProperty("prompt_eval_duration") Duration promptEvalDuration, @JsonProperty("eval_count") Integer evalCount, - @JsonProperty("eval_duration") Duration evalDuration) { + @JsonProperty("eval_duration") Duration evalDuration + ) { } /** diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaModel.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaModel.java index d92750664..14299301a 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaModel.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaModel.java @@ -21,6 +21,7 @@ import org.springframework.ai.model.ChatModelDescription; * Helper class for common Ollama models. * * @author Siarhei Blashuk + * @author Thomas Vitale * @since 0.8.1 */ public enum OllamaModel implements ChatModelDescription { @@ -35,11 +36,27 @@ public enum OllamaModel implements ChatModelDescription { */ LLAMA3("llama3"), + /** + * The 8B language model from Meta. + */ + LLAMA3_1("llama3.1"), + /** * The 7B parameters model */ MISTRAL("mistral"), + /** + * A 12B model with 128k context length, built by Mistral AI in collaboration with + * NVIDIA. + */ + MISTRAL_NEMO("mistral-nemo"), + + /** + * A small vision language model designed to run efficiently on edge devices. + */ + MOONDREAM("moondream"), + /** * The 2.7B uncensored Dolphin model */ diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaOptions.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaOptions.java index ebe6938b5..32c255073 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaOptions.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaOptions.java @@ -40,6 +40,7 @@ import com.fasterxml.jackson.annotation.JsonProperty; * Helper class for creating strongly-typed Ollama options. * * @author Christian Tzolov + * @author Thomas Vitale * @since 0.8.0 * @see Ollama @@ -53,11 +54,14 @@ public class OllamaOptions implements FunctionCallingOptions, ChatOptions, Embed private static final List NON_SUPPORTED_FIELDS = List.of("model", "format", "keep_alive"); - // Following fields are ptions which must be set when the model is loaded into memory. + // Following fields are options which must be set when the model is loaded into + // memory. + // See: https://github.com/ggerganov/llama.cpp/blob/master/examples/main/README.md // @formatter:off + /** - * useNUMA Whether to use NUMA. + * Whether to use NUMA. (Default: false) */ @JsonProperty("numa") private Boolean useNUMA; @@ -67,63 +71,78 @@ public class OllamaOptions implements FunctionCallingOptions, ChatOptions, Embed @JsonProperty("num_ctx") private Integer numCtx; /** - * ??? + * Prompt processing maximum batch size. (Default: 512) */ @JsonProperty("num_batch") private Integer numBatch; /** * The number of layers to send to the GPU(s). On macOS, it defaults to 1 * to enable metal support, 0 to disable. - */ + * (Default: -1, which indicates that numGPU should be set dynamically) + */ @JsonProperty("num_gpu") private Integer numGPU; /** - * ??? + * When using multiple GPUs this option controls which GPU is used + * for small tensors for which the overhead of splitting the computation + * across all GPUs is not worthwhile. The GPU in question will use slightly + * more VRAM to store a scratch buffer for temporary results. + * By default, GPU 0 is used. */ @JsonProperty("main_gpu")private Integer mainGPU; /** - * ??? + * (Default: false) */ @JsonProperty("low_vram") private Boolean lowVRAM; /** - * ??? + * (Default: true) */ @JsonProperty("f16_kv") private Boolean f16KV; /** - * ??? + * Return logits for all the tokens, not just the last one. + * To enable completions to return logprobs, this must be true. */ @JsonProperty("logits_all") private Boolean logitsAll; /** - * ??? + * Load only the vocabulary, not the weights. */ @JsonProperty("vocab_only") private Boolean vocabOnly; /** - * ??? + * By default, models are mapped into memory, which allows the system to load only the necessary parts + * of the model as needed. However, if the model is larger than your total amount of RAM or if your system is low + * on available memory, using mmap might increase the risk of pageouts, negatively impacting performance. + * Disabling mmap results in slower load times but may reduce pageouts if you're not using mlock. + * Note that if the model is larger than the total amount of RAM, turning off mmap would prevent + * the model from loading at all. + * (Default: null) */ @JsonProperty("use_mmap") private Boolean useMMap; /** - * ??? + * Lock the model in memory, preventing it from being swapped out when memory-mapped. + * This can improve performance but trades away some of the advantages of memory-mapping + * by requiring more RAM to run and potentially slowing down load times as the model loads into RAM. + * (Default: false) */ @JsonProperty("use_mlock") private Boolean useMLock; /** - * Sets the number of threads to use during computation. By default, - * Ollama will detect this for optimal performance. It is recommended to set this - * value to the number of physical CPU cores your system has (as opposed to the - * logical number of cores). + * Set the number of threads to use during generation. For optimal performance, it is recommended to set this value + * to the number of physical CPU cores your system has (as opposed to the logical number of cores). + * Using the correct number of threads can greatly improve performance. + * By default, Ollama will detect this value for optimal performance. */ @JsonProperty("num_thread") private Integer numThread; // Following fields are predict options used at runtime. /** - * ??? + * (Default: 4) */ @JsonProperty("num_keep") private Integer numKeep; @@ -162,7 +181,7 @@ public class OllamaOptions implements FunctionCallingOptions, ChatOptions, Embed @JsonProperty("tfs_z") private Float tfsZ; /** - * ??? + * (Default: 1.0) */ @JsonProperty("typical_p") private Float typicalP; @@ -186,12 +205,12 @@ public class OllamaOptions implements FunctionCallingOptions, ChatOptions, Embed @JsonProperty("repeat_penalty") private Float repeatPenalty; /** - * ??? + * (Default: 0.0) */ @JsonProperty("presence_penalty") private Float presencePenalty; /** - * ??? + * (Default: 0.0) */ @JsonProperty("frequency_penalty") private Float frequencyPenalty; @@ -215,7 +234,7 @@ public class OllamaOptions implements FunctionCallingOptions, ChatOptions, Embed @JsonProperty("mirostat_eta") private Float mirostatEta; /** - * ??? + * (Default: true) */ @JsonProperty("penalize_newline") private Boolean penalizeNewline; 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 6da2d04a0..6f7334d2c 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 @@ -36,6 +36,7 @@ import org.springframework.ai.chat.model.Generation; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.function.FunctionCallbackWrapper; 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.api.tool.MockWeatherService; import org.springframework.beans.factory.annotation.Autowired; @@ -55,10 +56,10 @@ class OllamaChatModelFunctionCallingIT { private static final Logger logger = LoggerFactory.getLogger(OllamaChatModelFunctionCallingIT.class); - private static String MODEL = "mistral"; + private static final String MODEL = OllamaModel.MISTRAL.getName(); @Container - static OllamaContainer ollamaContainer = new OllamaContainer("ollama/ollama:0.2.8"); + static OllamaContainer ollamaContainer = new OllamaContainer(OllamaImage.DEFAULT_IMAGE); static String baseUrl = "http://localhost:11434"; 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 dcbe44803..aea4daa1c 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 @@ -63,7 +63,7 @@ class OllamaChatModelIT { private static final Log logger = LogFactory.getLog(OllamaChatModelIT.class); @Container - static OllamaContainer ollamaContainer = new OllamaContainer("ollama/ollama:0.2.8"); + static OllamaContainer ollamaContainer = new OllamaContainer(OllamaImage.DEFAULT_IMAGE); static String baseUrl = "http://localhost:11434"; 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 c9bdeca09..c47755502 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 @@ -47,12 +47,12 @@ import static org.assertj.core.api.Assertions.assertThat; @Disabled("For manual smoke testing only.") class OllamaChatModelMultimodalIT { - private static final String MODEL = OllamaModel.MISTRAL.getName(); + private static final String MODEL = OllamaModel.MOONDREAM.getName(); private static final Log logger = LogFactory.getLog(OllamaChatModelIT.class); @Container - static OllamaContainer ollamaContainer = new OllamaContainer("ollama/ollama:0.2.8"); + static OllamaContainer ollamaContainer = new OllamaContainer(OllamaImage.DEFAULT_IMAGE); static String baseUrl = "http://localhost:11434"; 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 a616823fc..a10f0cef5 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 @@ -23,6 +23,7 @@ import org.apache.commons.logging.LogFactory; import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; +import org.springframework.ai.ollama.api.OllamaModel; import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.ollama.OllamaContainer; @@ -43,14 +44,14 @@ import static org.assertj.core.api.Assertions.assertThat; @Testcontainers class OllamaEmbeddingModelIT { - private static String MODEL = "llava"; + private static final String MODEL = OllamaModel.MISTRAL.getName(); private static final Log logger = LogFactory.getLog(OllamaApiIT.class); @Container - static OllamaContainer ollamaContainer = new OllamaContainer("ollama/ollama:0.1.32"); + static OllamaContainer ollamaContainer = new OllamaContainer(OllamaImage.DEFAULT_IMAGE); - static String baseUrl; + static String baseUrl = "http://localhost:11434"; @BeforeAll public static void beforeAll() throws IOException, InterruptedException { diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaImage.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaImage.java new file mode 100644 index 000000000..76d1a65ae --- /dev/null +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaImage.java @@ -0,0 +1,27 @@ +/* + * Copyright 2024 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.ai.ollama; + +import org.testcontainers.utility.DockerImageName; + +/** + * @author Thomas Vitale + */ +public class OllamaImage { + + public static final DockerImageName DEFAULT_IMAGE = DockerImageName.parse("ollama/ollama:0.2.8"); + +} 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 325470622..ced1a9d91 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 @@ -24,6 +24,7 @@ import org.apache.commons.logging.LogFactory; import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; +import org.springframework.ai.ollama.OllamaImage; import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.ollama.OllamaContainer; @@ -42,25 +43,31 @@ import static org.assertj.core.api.Assertions.assertThat;; /** * @author Christian Tzolov + * @author Thomas Vitale */ @Disabled("For manual smoke testing only.") @Testcontainers public class OllamaApiIT { + private static final String MODEL = OllamaModel.ORCA_MINI.getName(); + private static final Log logger = LogFactory.getLog(OllamaApiIT.class); @Container - static OllamaContainer ollamaContainer = new OllamaContainer("ollama/ollama:0.1.32"); + static OllamaContainer ollamaContainer = new OllamaContainer(OllamaImage.DEFAULT_IMAGE); static OllamaApi ollamaApi; + static String baseUrl = "http://localhost:11434"; + @BeforeAll public static void beforeAll() throws IOException, InterruptedException { - logger.info("Start pulling the 'orca-mini' generative (3GB) ... would take several minutes ..."); - ollamaContainer.execInContainer("ollama", "pull", "orca-mini"); - logger.info("orca-mini pulling competed!"); + logger.info("Start pulling the '" + MODEL + " ' generative ... would take several minutes ..."); + ollamaContainer.execInContainer("ollama", "pull", MODEL); + logger.info(MODEL + " pulling competed!"); - ollamaApi = new OllamaApi("http://" + ollamaContainer.getHost() + ":" + ollamaContainer.getMappedPort(11434)); + baseUrl = "http://" + ollamaContainer.getHost() + ":" + ollamaContainer.getMappedPort(11434); + ollamaApi = new OllamaApi(baseUrl); } @Test @@ -68,7 +75,7 @@ public class OllamaApiIT { var request = GenerateRequest .builder("What is the capital of Bulgaria and what is the size? What it the national anthem?") - .withModel("orca-mini") + .withModel(MODEL) .withStream(false) .build(); @@ -84,7 +91,7 @@ public class OllamaApiIT { @Test public void chat() { - var request = ChatRequest.builder("orca-mini") + var request = ChatRequest.builder(MODEL) .withStream(false) .withMessages(List.of( Message.builder(Role.SYSTEM) @@ -111,7 +118,7 @@ public class OllamaApiIT { @Test public void streamingChat() { - var request = ChatRequest.builder("orca-mini") + var request = ChatRequest.builder(MODEL) .withStream(true) .withMessages(List.of(Message.builder(Role.USER) .withContent("What is the capital of Bulgaria and what is the size? " + "What it the national anthem?") @@ -138,7 +145,7 @@ public class OllamaApiIT { @Test public void embedText() { - EmbeddingRequest request = new EmbeddingRequest("orca-mini", "I like to eat apples"); + EmbeddingRequest request = new EmbeddingRequest(MODEL, "I like to eat apples"); EmbeddingResponse response = ollamaApi.embeddings(request); 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 697a93c73..8c4a682e8 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 @@ -29,6 +29,7 @@ import org.junit.jupiter.api.Test; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.model.ModelOptionsUtils; +import org.springframework.ai.ollama.OllamaImage; import org.springframework.ai.ollama.api.OllamaApi; import org.springframework.ai.ollama.api.OllamaApi.ChatResponse; import org.springframework.ai.ollama.api.OllamaApi.Message; @@ -53,7 +54,7 @@ public class OllamaApiToolFunctionCallIT { MockWeatherService weatherService = new MockWeatherService(); @Container - static OllamaContainer ollamaContainer = new OllamaContainer("ollama/ollama:0.2.8"); + static OllamaContainer ollamaContainer = new OllamaContainer(OllamaImage.DEFAULT_IMAGE); static String baseUrl = "http://localhost:11434";