Add test for BeanOutputParser and support for JacksonAnnotations
Fixes #53
This commit is contained in:
committed by
Mark Pollack
parent
f9ca032cb3
commit
c15663137c
@@ -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 <T> The target type to convert the output into.
|
||||
* @param <T> The target type to which the output will be converted.
|
||||
* @author Mark Pollack
|
||||
* @author Christian Tzolov
|
||||
* @author Sebastian Ullrich
|
||||
*/
|
||||
public class BeanOutputParser<T> implements OutputParser<T> {
|
||||
|
||||
/** 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<T> 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<T> 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<T> implements OutputParser<T> {
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 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);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
Reference in New Issue
Block a user