diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatClient.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatClient.java index 8d8d76239..e1e9f640a 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatClient.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatClient.java @@ -132,13 +132,13 @@ public class OllamaChatClient implements ChatClient, StreamingChatClient { List ollamaMessages = prompt.getInstructions() .stream() .filter(message -> message.getMessageType() == MessageType.USER - || message.getMessageType() == MessageType.ASSISTANT) + || message.getMessageType() == MessageType.ASSISTANT + || message.getMessageType() == MessageType.SYSTEM) .map(m -> OllamaApi.Message.builder(toRole(m)).withContent(m.getContent()).build()) .toList(); // runtime options - Map promptOptions = objectToMap(prompt.getOptions()); - Map clientOptionsToUse = merge(promptOptions, this.clientOptions, HashMap.class); + Map clientOptionsToUse = merge(prompt.getOptions(), this.clientOptions, HashMap.class); return ChatRequest.builder(model) .withStream(stream) @@ -169,6 +169,9 @@ public class OllamaChatClient implements ChatClient, StreamingChatClient { } public static T merge(Object source, Object target, Class clazz) { + if (source == null) { + source = Map.of(); + } Map sourceMap = objectToMap(source); Map targetMap = objectToMap(target); diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatClientIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatClientIT.java index dfff01a45..7ef6b8bfe 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatClientIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatClientIT.java @@ -47,7 +47,7 @@ class OllamaChatClientIT { private static final Log logger = LogFactory.getLog(OllamaChatClientIT.class); @Container - static GenericContainer ollamaContainer = new GenericContainer<>("ollama/ollama:0.1.16").withExposedPorts(11434); + static GenericContainer ollamaContainer = new GenericContainer<>("ollama/ollama:0.1.21").withExposedPorts(11434); static String baseUrl; @@ -86,7 +86,6 @@ class OllamaChatClientIT { assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); } - @Disabled("TODO: Fix the parser instructions to return the correct format") @Test void outputParser() { DefaultConversionService conversionService = new DefaultConversionService(); @@ -106,7 +105,6 @@ class OllamaChatClientIT { assertThat(list).hasSize(5); } - @Disabled("TODO: Fix the parser instructions to return the correct format") @Test void mapOutputParser() { MapOutputParser outputParser = new MapOutputParser(); @@ -131,7 +129,6 @@ class OllamaChatClientIT { record ActorsFilmsRecord(String actor, List movies) { } - @Disabled("TODO: Fix the parser instructions to return the correct format") @Test void beanOutputParserRecords() { @@ -141,7 +138,6 @@ class OllamaChatClientIT { String template = """ Generate the filmography of 5 movies for Tom Hanks. {format} - Remove Markdown code blocks from the output. """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); @@ -152,7 +148,6 @@ class OllamaChatClientIT { assertThat(actorsFilms.movies()).hasSize(5); } - @Disabled("TODO: Fix the parser instructions to return the correct format") @Test void beanStreamOutputParserRecords() { diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingClientIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingClientIT.java index 50e8621cd..6eb19a801 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingClientIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingClientIT.java @@ -30,7 +30,7 @@ class OllamaEmbeddingClientIT { private static final Log logger = LogFactory.getLog(OllamaApiIT.class); @Container - static GenericContainer ollamaContainer = new GenericContainer<>("ollama/ollama:0.1.16").withExposedPorts(11434); + static GenericContainer ollamaContainer = new GenericContainer<>("ollama/ollama:0.1.21").withExposedPorts(11434); static String baseUrl; 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 da34fa839..db8dbb108 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 @@ -51,7 +51,7 @@ public class OllamaApiIT { private static final Log logger = LogFactory.getLog(OllamaApiIT.class); @Container - static GenericContainer ollamaContainer = new GenericContainer<>("ollama/ollama:0.1.16").withExposedPorts(11434); + static GenericContainer ollamaContainer = new GenericContainer<>("ollama/ollama:0.1.21").withExposedPorts(11434); static OllamaApi ollamaApi; @@ -87,9 +87,14 @@ public class OllamaApiIT { var request = ChatRequest.builder("orca-mini") .withStream(false) - .withMessages(List.of(Message.builder(Role.USER) - .withContent("What is the capital of Bulgaria and what is the size? " + "What it the national anthem?") - .build())) + .withMessages(List.of( + Message.builder(Role.SYSTEM) + .withContent("You are geography teacher. You are talking to a student.") + .build(), + Message.builder(Role.USER) + .withContent("What is the capital of Bulgaria and what is the size? " + + "What it the national anthem?") + .build())) .withOptions(OllamaOptions.create().withTemperature(0.9f)) .build(); @@ -127,7 +132,7 @@ public class OllamaApiIT { .collect(Collectors.joining("\n"))).contains("Sofia"); ChatResponse lastResponse = responses.get(responses.size() - 1); - assertThat(lastResponse.message()).isNull(); + assertThat(lastResponse.message().content()).isEmpty(); assertThat(lastResponse.done()).isTrue(); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/parser/BeanOutputParser.java b/spring-ai-core/src/main/java/org/springframework/ai/parser/BeanOutputParser.java index 718153a54..8c8f4326a 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/parser/BeanOutputParser.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/parser/BeanOutputParser.java @@ -25,6 +25,7 @@ import com.github.victools.jsonschema.generator.SchemaGeneratorConfig; import com.github.victools.jsonschema.generator.SchemaGeneratorConfigBuilder; import com.github.victools.jsonschema.module.jackson.JacksonModule; +import java.util.Map; import java.util.Objects; import static com.github.victools.jsonschema.generator.OptionPreset.PLAIN_JSON; @@ -95,6 +96,13 @@ public class BeanOutputParser implements OutputParser { */ public T parse(String text) { try { + // If the response is a JSON Schema, extract the properties and use them as + // the + // response. + Map map = this.objectMapper.readValue(text, Map.class); + if (map.containsKey("$schema")) { + text = this.objectMapper.writeValueAsString(map.get("properties")); + } return (T) this.objectMapper.readValue(text, this.clazz); } catch (JsonProcessingException e) { @@ -122,6 +130,7 @@ public class BeanOutputParser implements OutputParser { String template = """ Your response should be in JSON format. Do not include any explanations, only provide a RFC8259 compliant JSON response following this format without deviation. + Do not include markdown code blocks in your response. Here is the JSON Schema instance your output must adhere to: ```%s``` """;