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 919c4f456..5364aa0d2 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 @@ -20,59 +20,82 @@ import com.fasterxml.jackson.core.JsonProcessingException; import com.fasterxml.jackson.databind.DeserializationFeature; import com.fasterxml.jackson.databind.JsonNode; import com.fasterxml.jackson.databind.ObjectMapper; -import com.github.victools.jsonschema.generator.OptionPreset; import com.github.victools.jsonschema.generator.SchemaGenerator; import com.github.victools.jsonschema.generator.SchemaGeneratorConfig; import com.github.victools.jsonschema.generator.SchemaGeneratorConfigBuilder; -import com.github.victools.jsonschema.generator.SchemaVersion; -import org.springframework.ai.prompt.PromptTemplate; +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; +import static com.github.victools.jsonschema.generator.SchemaVersion.DRAFT_2020_12; + /** - * {@link OutputParser} implementation that uses JSON schema to convert the LLM output - * into a desired object of type T. + * An implementation of {@link OutputParser} that transforms the LLM output to a specific + * object type using JSON schema. This parser works by generating a JSON schema based on a + * given Java class, which is then used to validate and transform the LLM output into the + * desired type. * - * @param The target type to convert the output into. + * @param The target type to which the output will be converted. * @author Mark Pollack * @author Christian Tzolov + * @author Sebastian Ullrich */ public class BeanOutputParser implements OutputParser { + /** Holds the generated JSON schema for the target type. */ private String jsonSchema; + /** The Java class representing the target type. */ + @SuppressWarnings({ "FieldMayBeFinal", "rawtypes" }) private Class clazz; + /** The object mapper used for deserialization and other JSON operations. */ + @SuppressWarnings("FieldMayBeFinal") private ObjectMapper objectMapper; + /** + * Constructor to initialize with the target type's class. + * @param clazz The target type's class. + */ public BeanOutputParser(Class clazz) { - Objects.requireNonNull(clazz, "Java Class can not be null;"); - this.clazz = clazz; - this.objectMapper = getObjectMapper(); - generateSchema(); + this(clazz, null); } + /** + * Constructor to initialize with the target type's class and a custom object mapper. + * @param clazz The target type's class. + * @param objectMapper Custom object mapper for JSON operations. + */ public BeanOutputParser(Class clazz, ObjectMapper objectMapper) { - Objects.requireNonNull(clazz, "Java Class can not be null;"); - Objects.requireNonNull(objectMapper, "ObjectMapper can not be null;"); + Objects.requireNonNull(clazz, "Java Class cannot be null;"); this.clazz = clazz; - this.objectMapper = objectMapper; + this.objectMapper = objectMapper != null ? objectMapper : getObjectMapper(); generateSchema(); } + /** + * Generates the JSON schema for the target type. + */ private void generateSchema() { - SchemaGeneratorConfigBuilder configBuilder = new SchemaGeneratorConfigBuilder(SchemaVersion.DRAFT_2020_12, - OptionPreset.PLAIN_JSON); + JacksonModule jacksonModule = new JacksonModule(); + SchemaGeneratorConfigBuilder configBuilder = new SchemaGeneratorConfigBuilder(DRAFT_2020_12, PLAIN_JSON) + .with(jacksonModule); SchemaGeneratorConfig config = configBuilder.build(); SchemaGenerator generator = new SchemaGenerator(config); - JsonNode jsonSchema = generator.generateSchema(this.clazz); - this.jsonSchema = jsonSchema.toPrettyString(); + JsonNode jsonNode = generator.generateSchema(this.clazz); + this.jsonSchema = jsonNode.toPrettyString(); } @Override + /** + * Parses the given text to transform it to the desired target type. + * @param text The LLM output in string format. + * @return The parsed output in the desired target type. + */ public T parse(String text) { try { + // noinspection unchecked return (T) this.objectMapper.readValue(text, this.clazz); } catch (JsonProcessingException e) { @@ -80,23 +103,30 @@ public class BeanOutputParser implements OutputParser { } } + /** + * Configures and returns an object mapper for JSON operations. + * @return Configured object mapper. + */ protected ObjectMapper getObjectMapper() { - ObjectMapper objectMapper = new ObjectMapper(); - objectMapper.configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, false); - return objectMapper; + ObjectMapper mapper = new ObjectMapper(); + mapper.configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, false); + return mapper; } + /** + * Provides the expected format of the response, instructing that it should adhere to + * the generated JSON schema. + * @return The instruction format string. + */ @Override public String getFormat() { - - String raw = """ + 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. Here is the JSON Schema instance your output must adhere to: - ```{jsonSchema}``` + ```%s``` """; - PromptTemplate promptTemplate = new PromptTemplate(raw); - return promptTemplate.render(Map.of("jsonSchema", this.jsonSchema)); + return String.format(template, this.jsonSchema); } } diff --git a/spring-ai-core/src/test/java/org/springframework/ai/parser/BeanOutputParserTest.java b/spring-ai-core/src/test/java/org/springframework/ai/parser/BeanOutputParserTest.java new file mode 100644 index 000000000..28c47e0a7 --- /dev/null +++ b/spring-ai-core/src/test/java/org/springframework/ai/parser/BeanOutputParserTest.java @@ -0,0 +1,137 @@ +package org.springframework.ai.parser; + +import com.fasterxml.jackson.annotation.JsonProperty; +import com.fasterxml.jackson.annotation.JsonPropertyDescription; +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.DeserializationFeature; +import com.fasterxml.jackson.databind.ObjectMapper; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.when; + +/** + * @author Sebastian Ullrich + */ +@ExtendWith(MockitoExtension.class) +class BeanOutputParserTest { + + @Mock + private ObjectMapper objectMapperMock; + + @Test + public void shouldHavePreconfiguredDefaultObjectMapper() { + var parser = new BeanOutputParser<>(TestClass.class); + var objectMapper = parser.getObjectMapper(); + assertThat(objectMapper.isEnabled(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES)).isFalse(); + } + + @Test + public void shouldUseProvidedObjectMapperForParsing() throws JsonProcessingException { + var testClass = new TestClass("some string"); + when(objectMapperMock.readValue(anyString(), eq(TestClass.class))).thenReturn(testClass); + var parser = new BeanOutputParser<>(TestClass.class, objectMapperMock); + assertThat(parser.parse("{}")).isEqualTo(testClass); + } + + @Nested + class ParserTest { + + @Test + public void shouldParseFieldNamesFromString() { + var parser = new BeanOutputParser<>(TestClass.class); + var testClass = parser.parse("{ \"someString\": \"some value\" }"); + assertThat(testClass.getSomeString()).isEqualTo("some value"); + } + + @Test + public void shouldParseJsonPropertiesFromString() { + var parser = new BeanOutputParser<>(TestClassWithJsonAnnotations.class); + var testClass = parser.parse("{ \"string_property\": \"some value\" }"); + assertThat(testClass.getSomeString()).isEqualTo("some value"); + } + + } + + @Nested + class FormatTest { + + @Test + public void shouldReturnFormatContainingResponseInstructionsAndJsonSchema() { + var parser = new BeanOutputParser<>(TestClass.class); + assertThat(parser.getFormat()).isEqualTo( + """ + Your response should be in JSON format. + Do not include any explanations, only provide a RFC8259 compliant JSON response following this format without deviation. + Here is the JSON Schema instance your output must adhere to: + ```{ + "$schema" : "https://json-schema.org/draft/2020-12/schema", + "type" : "object", + "properties" : { + "someString" : { + "type" : "string" + } + } + }``` + """); + } + + @Test + public void shouldReturnFormatContainingJsonSchemaIncludingPropertyAndPropertyDescription() { + var parser = new BeanOutputParser<>(TestClassWithJsonAnnotations.class); + assertThat(parser.getFormat()).contains(""" + ```{ + "$schema" : "https://json-schema.org/draft/2020-12/schema", + "type" : "object", + "properties" : { + "string_property" : { + "type" : "string", + "description" : "string_property_description" + } + } + }``` + """); + } + + } + + public static class TestClass { + + private String someString; + + @SuppressWarnings("unused") + public TestClass() { + } + + public TestClass(String someString) { + this.someString = someString; + } + + public String getSomeString() { + return someString; + } + + } + + public static class TestClassWithJsonAnnotations { + + @JsonProperty("string_property") + @JsonPropertyDescription("string_property_description") + private String someString; + + public TestClassWithJsonAnnotations() { + } + + public String getSomeString() { + return someString; + } + + } + +}