diff --git a/wayang-commons/wayang-basic/src/main/java/org/apache/wayang/basic/function/ProjectionDescriptor.java b/wayang-commons/wayang-basic/src/main/java/org/apache/wayang/basic/function/ProjectionDescriptor.java index 15aeabc53..f8a36555e 100644 --- a/wayang-commons/wayang-basic/src/main/java/org/apache/wayang/basic/function/ProjectionDescriptor.java +++ b/wayang-commons/wayang-basic/src/main/java/org/apache/wayang/basic/function/ProjectionDescriptor.java @@ -45,7 +45,7 @@ private static class PojoImplementation private final String fieldName; - private Field field; + private transient Field field; private PojoImplementation(final String fieldName) { this.fieldName = fieldName; @@ -54,6 +54,10 @@ private PojoImplementation(final String fieldName) { @Override @SuppressWarnings("unchecked") public Output apply(final Input input) { + if (input == null) { + return null; + } + // Initialization code. if (this.field == null) { @@ -127,7 +131,7 @@ public static ProjectionDescriptor createForRecords(final Record } private static FunctionDescriptor.SerializableFunction createPojoJavaImplementation( - final String[] fieldNames, final BasicDataUnitType inputType) { + final String[] fieldNames) { // Get the names of the fields to be projected. if (fieldNames.length != 1) { return t -> { @@ -185,7 +189,7 @@ public ProjectionDescriptor(final Class inputTypeClass, */ public ProjectionDescriptor(final BasicDataUnitType inputType, final BasicDataUnitType outputType, final String... fieldNames) { - this(createPojoJavaImplementation(fieldNames, inputType), + this(createPojoJavaImplementation(fieldNames), Collections.unmodifiableList(Arrays.asList(fieldNames)), inputType, outputType); diff --git a/wayang-commons/wayang-basic/src/test/java/org/apache/wayang/basic/function/ProjectionDescriptorTest.java b/wayang-commons/wayang-basic/src/test/java/org/apache/wayang/basic/function/ProjectionDescriptorTest.java index 72f60994d..31ffa5812 100644 --- a/wayang-commons/wayang-basic/src/test/java/org/apache/wayang/basic/function/ProjectionDescriptorTest.java +++ b/wayang-commons/wayang-basic/src/test/java/org/apache/wayang/basic/function/ProjectionDescriptorTest.java @@ -25,7 +25,14 @@ import java.util.function.Function; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertThrows; + +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.ObjectInputStream; +import java.io.ObjectOutputStream; /** * Tests for the {@link ProjectionDescriptor}. @@ -45,11 +52,60 @@ void testPojoImplementation() { stringImplementation.apply(new Pojo("testValue", 1)) ); assertNull(stringImplementation.apply(new Pojo(null, 1))); + assertNull(stringImplementation.apply(null)); assertEquals( Integer.valueOf(1), integerImplementation.apply(new Pojo("testValue", 1)) ); + } + + @Test + void testPojoImplementationMultipleFieldsThrows() { + final ProjectionDescriptor multiFieldDescriptor = + new ProjectionDescriptor<>(Pojo.class, String.class, "string", "integer"); + final Function multiFieldImplementation = multiFieldDescriptor.getJavaImplementation(); + + assertThrows( + IllegalStateException.class, + () -> multiFieldImplementation.apply(new Pojo("testValue", 1)) + ); + } + + @Test + void testPojoImplementationNonExistentFieldThrows() { + final ProjectionDescriptor invalidDescriptor = + new ProjectionDescriptor<>(Pojo.class, String.class, "nonExistentField"); + final Function invalidImplementation = invalidDescriptor.getJavaImplementation(); + + assertThrows( + IllegalStateException.class, + () -> invalidImplementation.apply(new Pojo("testValue", 1)) + ); + } + + @Test + @SuppressWarnings("unchecked") + void testPojoImplementationSerialization() throws Exception { + final ProjectionDescriptor stringDescriptor = + new ProjectionDescriptor<>(Pojo.class, String.class, "string"); + Function fn = stringDescriptor.getJavaImplementation(); + + // Use the function once so that 'field' is populated + assertEquals("val1", fn.apply(new Pojo("val1", 42))); + + // Serialize and deserialize + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + try (ObjectOutputStream oos = new ObjectOutputStream(baos)) { + oos.writeObject(fn); + } + + Function deserializedFn; + try (ObjectInputStream ois = new ObjectInputStream(new ByteArrayInputStream(baos.toByteArray()))) { + deserializedFn = (Function) ois.readObject(); + } + assertNotNull(deserializedFn); + assertEquals("val2", deserializedFn.apply(new Pojo("val2", 99))); } @Test @@ -65,7 +121,7 @@ void testRecordImplementation() { ); } - public static class Pojo { + public static class Pojo implements java.io.Serializable { public String string;