diff --git a/src/main/java/it/cnr/isti/workflow/manager/llms/providers/LLMProvider.java b/src/main/java/it/cnr/isti/workflow/manager/llms/providers/LLMProvider.java index c986e35..45b796f 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/llms/providers/LLMProvider.java +++ b/src/main/java/it/cnr/isti/workflow/manager/llms/providers/LLMProvider.java @@ -11,6 +11,10 @@ public interface LLMProvider { List getRegisteredModels(); String generate(String model, String prompt); + default String generateJson(String model, String prompt) { + return generate(model, prompt); + } + default String generate(String model, String prompt, String authorization) { return generate(model, prompt); } diff --git a/src/main/java/it/cnr/isti/workflow/manager/llms/providers/ollama/InternalOllamaLLMProvider.java b/src/main/java/it/cnr/isti/workflow/manager/llms/providers/ollama/InternalOllamaLLMProvider.java index aed1151..64f08f4 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/llms/providers/ollama/InternalOllamaLLMProvider.java +++ b/src/main/java/it/cnr/isti/workflow/manager/llms/providers/ollama/InternalOllamaLLMProvider.java @@ -1,6 +1,7 @@ package it.cnr.isti.workflow.manager.llms.providers.ollama; import java.time.Duration; +import java.util.LinkedHashMap; import java.util.List; import java.util.Map; import java.util.Objects; @@ -46,13 +47,28 @@ public class InternalOllamaLLMProvider implements LLMProvider { } public String generate(String model, String prompt) { + return generate(model, prompt, false); + } + + @Override + public String generateJson(String model, String prompt) { + return generate(model, prompt, true); + } + + private String generate(String model, String prompt, boolean jsonResponse) { Objects.requireNonNull(prompt, "prompt cannot be null"); Objects.requireNonNull(model, "model cannot be null"); ObjectMapper mapper = new ObjectMapper(); - Map bodyMap = Map.of( - "model", model, - "prompt", prompt, - "stream", false); + Map bodyMap = new LinkedHashMap<>(); + bodyMap.put("model", model); + bodyMap.put("prompt", prompt); + bodyMap.put("stream", false); + if (jsonResponse) { + bodyMap.put("format", "json"); + bodyMap.put("options", Map.of( + "temperature", 0.1, + "num_predict", 4096)); + } // Implement the logic to call the Ollama API and return the response WebClient webClient = webClientBuilder.baseUrl(this.ollamaURL).build();