diff --git a/sdks/java/core/src/main/java/org/apache/beam/sdk/io/TFRecordWriteSchemaTransformProvider.java b/sdks/java/core/src/main/java/org/apache/beam/sdk/io/TFRecordWriteSchemaTransformProvider.java index bc9b7bbeac66..c787997b1e77 100644 --- a/sdks/java/core/src/main/java/org/apache/beam/sdk/io/TFRecordWriteSchemaTransformProvider.java +++ b/sdks/java/core/src/main/java/org/apache/beam/sdk/io/TFRecordWriteSchemaTransformProvider.java @@ -184,8 +184,7 @@ public PCollectionRowTuple expand(PCollectionRowTuple input) { output = ""; } } - PCollection errorOutput = - byteArrays.get(ERROR_TAG).setRowSchema(ErrorHandling.errorSchema(errorSchema)); + PCollection errorOutput = byteArrays.get(ERROR_TAG).setRowSchema(errorSchema); return PCollectionRowTuple.of(handleErrors ? output : "errors", errorOutput); } } diff --git a/sdks/java/core/src/test/java/org/apache/beam/sdk/io/TFRecordSchemaTransformProviderTest.java b/sdks/java/core/src/test/java/org/apache/beam/sdk/io/TFRecordSchemaTransformProviderTest.java index 9c067a533e0b..65e38ed9d298 100644 --- a/sdks/java/core/src/test/java/org/apache/beam/sdk/io/TFRecordSchemaTransformProviderTest.java +++ b/sdks/java/core/src/test/java/org/apache/beam/sdk/io/TFRecordSchemaTransformProviderTest.java @@ -49,6 +49,7 @@ import org.apache.beam.sdk.schemas.Schema; import org.apache.beam.sdk.schemas.transforms.SchemaTransform; import org.apache.beam.sdk.schemas.transforms.SchemaTransformProvider; +import org.apache.beam.sdk.schemas.transforms.providers.ErrorHandling; import org.apache.beam.sdk.testing.NeedsRunner; import org.apache.beam.sdk.testing.PAssert; import org.apache.beam.sdk.testing.TestPipeline; @@ -238,6 +239,34 @@ public void testWriteBuildTransform() { .build()); } + @Test + public void testWriteErrorSchemaMatchesTheRowsErrorFnEmits() throws Exception { + // The schema is fixed when the graph is built, so this needs no runner. + writePipeline.enableAbandonedNodeEnforcement(false); + + Schema schema = Schema.of(Schema.Field.of("record", Schema.FieldType.BYTES)); + + TFRecordWriteSchemaTransformProvider provider = new TFRecordWriteSchemaTransformProvider(); + TFRecordWriteSchemaTransform transform = + (TFRecordWriteSchemaTransform) + provider.from( + TFRecordWriteSchemaTransformConfiguration.builder() + .setOutputPrefix(tempFolder.getRoot().toPath().resolve("errors").toString()) + .setCompression("UNCOMPRESSED") + .setNumShards(0) + .setNoSpilling(true) + .build()); + + Row row = Row.withSchema(schema).addValue("foo".getBytes(StandardCharsets.UTF_8)).build(); + PCollection input = + writePipeline.apply(Create.of(Collections.singletonList(row)).withRowSchema(schema)); + PCollectionRowTuple result = PCollectionRowTuple.of("input", input).apply(transform); + + // ErrorFn emits ErrorHandling.errorRecord(errorSchema, ..), so the collection has to carry + // that schema. Wrapping it a second time declares a shape no element it emits can match. + assertEquals(ErrorHandling.errorSchema(schema), result.get("errors").getSchema()); + } + @Test public void testReadFindTransformAndMakeItWork() { ServiceLoader serviceLoader =