This is an automated email from the ASF dual-hosted git repository.
Abacn pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/beam.git
The following commit(s) were added to refs/heads/master by this push:
new 7de4d3ab73b Fix KafkaIO wrtie SchemaTransform parallelism (#39844)
7de4d3ab73b is described below
commit 7de4d3ab73b6951178f6eb625b55597b2ee9081b
Author: Yi Hu <[email protected]>
AuthorDate: Fri Aug 21 16:32:54 2026 -0400
Fix KafkaIO wrtie SchemaTransform parallelism (#39844)
* Fix KafkaIO wrtie SchemaTransform parallelism
* Clean up redundant annotations
---
.../kafka/KafkaWriteSchemaTransformProvider.java | 41 ++++++++++------------
.../KafkaWriteSchemaTransformProviderTest.java | 27 +++++++++-----
2 files changed, 36 insertions(+), 32 deletions(-)
diff --git
a/sdks/java/io/kafka/src/main/java/org/apache/beam/sdk/io/kafka/KafkaWriteSchemaTransformProvider.java
b/sdks/java/io/kafka/src/main/java/org/apache/beam/sdk/io/kafka/KafkaWriteSchemaTransformProvider.java
index b9c41746240..d8d0fe478e3 100644
---
a/sdks/java/io/kafka/src/main/java/org/apache/beam/sdk/io/kafka/KafkaWriteSchemaTransformProvider.java
+++
b/sdks/java/io/kafka/src/main/java/org/apache/beam/sdk/io/kafka/KafkaWriteSchemaTransformProvider.java
@@ -28,11 +28,11 @@ import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.Set;
-import javax.annotation.Nullable;
import org.apache.avro.generic.GenericRecord;
import org.apache.beam.model.pipeline.v1.ExternalTransforms;
import org.apache.beam.sdk.coders.ByteArrayCoder;
import org.apache.beam.sdk.coders.KvCoder;
+import org.apache.beam.sdk.coders.NullableCoder;
import org.apache.beam.sdk.extensions.avro.coders.AvroCoder;
import org.apache.beam.sdk.extensions.avro.schemas.utils.AvroUtils;
import org.apache.beam.sdk.extensions.protobuf.ProtoByteUtils;
@@ -63,9 +63,7 @@ import org.apache.beam.sdk.values.TupleTag;
import org.apache.beam.sdk.values.TupleTagList;
import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.Sets;
import org.apache.kafka.common.serialization.ByteArraySerializer;
-import org.checkerframework.checker.initialization.qual.Initialized;
-import org.checkerframework.checker.nullness.qual.NonNull;
-import org.checkerframework.checker.nullness.qual.UnknownKeyFor;
+import org.checkerframework.checker.nullness.qual.Nullable;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
@@ -78,22 +76,20 @@ public class KafkaWriteSchemaTransformProvider
public static final Set<String> SUPPORTED_FORMATS =
Sets.newHashSet(SUPPORTED_FORMATS_STR.split(","));
public static final TupleTag<Row> ERROR_TAG = new TupleTag<Row>() {};
- public static final TupleTag<KV<byte[], byte[]>> OUTPUT_TAG =
- new TupleTag<KV<byte[], byte[]>>() {};
- public static final TupleTag<KV<byte[], GenericRecord>> RECORD_OUTPUT_TAG =
- new TupleTag<KV<byte[], GenericRecord>>() {};
+ public static final TupleTag<KV<byte @Nullable [], byte[]>> OUTPUT_TAG =
+ new TupleTag<KV<byte @Nullable [], byte[]>>() {};
+ public static final TupleTag<KV<byte @Nullable [], GenericRecord>>
RECORD_OUTPUT_TAG =
+ new TupleTag<KV<byte @Nullable [], GenericRecord>>() {};
private static final Logger LOG =
LoggerFactory.getLogger(KafkaWriteSchemaTransformProvider.class);
@Override
- protected @UnknownKeyFor @NonNull @Initialized
Class<KafkaWriteSchemaTransformConfiguration>
- configurationClass() {
+ protected Class<KafkaWriteSchemaTransformConfiguration> configurationClass()
{
return KafkaWriteSchemaTransformConfiguration.class;
}
@Override
- protected @UnknownKeyFor @NonNull @Initialized SchemaTransform from(
- KafkaWriteSchemaTransformConfiguration configuration) {
+ protected SchemaTransform from(KafkaWriteSchemaTransformConfiguration
configuration) {
if (!SUPPORTED_FORMATS.contains(configuration.getFormat())) {
throw new IllegalArgumentException(
"Format "
@@ -126,20 +122,20 @@ public class KafkaWriteSchemaTransformProvider
}
}
- public abstract static class BaseKafkaWriterFn<T> extends DoFn<Row,
KV<byte[], T>> {
+ public abstract static class BaseKafkaWriterFn<T> extends DoFn<Row,
KV<byte @Nullable [], T>> {
private final SerializableFunction<Row, T> conversionFn;
private final Counter errorCounter;
private Long errorsInBundle = 0L;
private final boolean handleErrors;
private final Schema errorSchema;
- private final TupleTag<KV<byte[], T>> successTag;
+ private final TupleTag<KV<byte @Nullable [], T>> successTag;
public BaseKafkaWriterFn(
String name,
SerializableFunction<Row, T> conversionFn,
Schema errorSchema,
boolean handleErrors,
- TupleTag<KV<byte[], T>> successTag) {
+ TupleTag<KV<byte @Nullable [], T>> successTag) {
this.conversionFn = conversionFn;
this.errorCounter =
Metrics.counter(KafkaWriteSchemaTransformProvider.class, name);
this.handleErrors = handleErrors;
@@ -149,9 +145,9 @@ public class KafkaWriteSchemaTransformProvider
@ProcessElement
public void process(@DoFn.Element Row row, MultiOutputReceiver receiver)
{
- KV<byte[], T> output = null;
+ KV<byte @Nullable [], T> output = null;
try {
- output = KV.of(new byte[1], conversionFn.apply(row));
+ output = KV.of(null, conversionFn.apply(row));
} catch (Exception e) {
if (!handleErrors) {
throw new RuntimeException(e);
@@ -262,7 +258,7 @@ public class KafkaWriteSchemaTransformProvider
HashMap<String, Object> producerConfig = new
HashMap<>(configOverrides);
outputTuple
.get(RECORD_OUTPUT_TAG)
- .setCoder(KvCoder.of(ByteArrayCoder.of(),
AvroCoder.of(avroSchema)))
+ .setCoder(KvCoder.of(NullableCoder.of(ByteArrayCoder.of()),
AvroCoder.of(avroSchema)))
.apply(
"Map Rows to GenericRecords",
KafkaIO.<byte[], GenericRecord>write()
@@ -284,6 +280,7 @@ public class KafkaWriteSchemaTransformProvider
outputTuple
.get(OUTPUT_TAG)
+ .setCoder(KvCoder.of(NullableCoder.of(ByteArrayCoder.of()),
ByteArrayCoder.of()))
.apply(
KafkaIO.<byte[], byte[]>write()
.withTopic(configuration.getTopic())
@@ -318,19 +315,17 @@ public class KafkaWriteSchemaTransformProvider
}
@Override
- public @UnknownKeyFor @NonNull @Initialized String identifier() {
+ public String identifier() {
return getUrn(ExternalTransforms.ManagedTransforms.Urns.KAFKA_WRITE);
}
@Override
- public @UnknownKeyFor @NonNull @Initialized List<@UnknownKeyFor @NonNull
@Initialized String>
- inputCollectionNames() {
+ public List<String> inputCollectionNames() {
return Collections.singletonList("input");
}
@Override
- public @UnknownKeyFor @NonNull @Initialized List<@UnknownKeyFor @NonNull
@Initialized String>
- outputCollectionNames() {
+ public List<String> outputCollectionNames() {
return Collections.emptyList();
}
diff --git
a/sdks/java/io/kafka/src/test/java/org/apache/beam/sdk/io/kafka/KafkaWriteSchemaTransformProviderTest.java
b/sdks/java/io/kafka/src/test/java/org/apache/beam/sdk/io/kafka/KafkaWriteSchemaTransformProviderTest.java
index fed783f70a0..cc5bd82a5a8 100644
---
a/sdks/java/io/kafka/src/test/java/org/apache/beam/sdk/io/kafka/KafkaWriteSchemaTransformProviderTest.java
+++
b/sdks/java/io/kafka/src/test/java/org/apache/beam/sdk/io/kafka/KafkaWriteSchemaTransformProviderTest.java
@@ -30,6 +30,7 @@ import org.apache.avro.generic.GenericRecord;
import org.apache.beam.sdk.Pipeline;
import org.apache.beam.sdk.coders.ByteArrayCoder;
import org.apache.beam.sdk.coders.KvCoder;
+import org.apache.beam.sdk.coders.NullableCoder;
import org.apache.beam.sdk.extensions.avro.coders.AvroCoder;
import org.apache.beam.sdk.extensions.avro.schemas.utils.AvroUtils;
import org.apache.beam.sdk.extensions.protobuf.ProtoByteUtils;
@@ -146,9 +147,9 @@ public class KafkaWriteSchemaTransformProviderTest {
public void testKafkaErrorFnSuccess() throws Exception {
List<KV<byte[], byte[]>> msg =
Arrays.asList(
- KV.of(new byte[1], "{\"name\":\"a\"}".getBytes(UTF_8)),
- KV.of(new byte[1], "{\"name\":\"b\"}".getBytes(UTF_8)),
- KV.of(new byte[1], "{\"name\":\"c\"}".getBytes(UTF_8)));
+ KV.of(null, "{\"name\":\"a\"}".getBytes(UTF_8)),
+ KV.of(null, "{\"name\":\"b\"}".getBytes(UTF_8)),
+ KV.of(null, "{\"name\":\"c\"}".getBytes(UTF_8)));
PCollection<Row> input = p.apply(Create.of(ROWS));
Schema errorSchema = ErrorHandling.errorSchema(BEAMSCHEMA);
@@ -159,6 +160,9 @@ public class KafkaWriteSchemaTransformProviderTest {
.withOutputTags(OUTPUT_TAG, TupleTagList.of(ERROR_TAG)));
output.get(ERROR_TAG).setRowSchema(errorSchema);
+ output
+ .get(OUTPUT_TAG)
+ .setCoder(KvCoder.of(NullableCoder.of(ByteArrayCoder.of()),
ByteArrayCoder.of()));
PAssert.that(output.get(OUTPUT_TAG)).containsInAnyOrder(msg);
p.run().waitUntilFinish();
@@ -168,9 +172,9 @@ public class KafkaWriteSchemaTransformProviderTest {
public void testKafkaErrorFnRawSuccess() throws Exception {
List<KV<byte[], byte[]>> msg =
Arrays.asList(
- KV.of(new byte[1], "a".getBytes(UTF_8)),
- KV.of(new byte[1], "b".getBytes(UTF_8)),
- KV.of(new byte[1], "c".getBytes(UTF_8)));
+ KV.of(null, "a".getBytes(UTF_8)),
+ KV.of(null, "b".getBytes(UTF_8)),
+ KV.of(null, "c".getBytes(UTF_8)));
PCollection<Row> input = p.apply(Create.of(RAW_ROWS));
Schema errorSchema = ErrorHandling.errorSchema(BEAM_RAW_SCHEMA);
@@ -182,6 +186,9 @@ public class KafkaWriteSchemaTransformProviderTest {
.withOutputTags(OUTPUT_TAG, TupleTagList.of(ERROR_TAG)));
output.get(ERROR_TAG).setRowSchema(errorSchema);
+ output
+ .get(OUTPUT_TAG)
+ .setCoder(KvCoder.of(NullableCoder.of(ByteArrayCoder.of()),
ByteArrayCoder.of()));
PAssert.that(output.get(OUTPUT_TAG)).containsInAnyOrder(msg);
p.run().waitUntilFinish();
@@ -199,6 +206,9 @@ public class KafkaWriteSchemaTransformProviderTest {
.withOutputTags(OUTPUT_TAG, TupleTagList.of(ERROR_TAG)));
output.get(ERROR_TAG).setRowSchema(errorSchema);
+ output
+ .get(OUTPUT_TAG)
+ .setCoder(KvCoder.of(NullableCoder.of(ByteArrayCoder.of()),
ByteArrayCoder.of()));
p.run().waitUntilFinish();
}
@@ -223,8 +233,7 @@ public class KafkaWriteSchemaTransformProviderTest {
record3.put("name", "c");
List<KV<byte[], GenericRecord>> msg =
- Arrays.asList(
- KV.of(new byte[1], record1), KV.of(new byte[1], record2),
KV.of(new byte[1], record3));
+ Arrays.asList(KV.of(null, record1), KV.of(null, record2), KV.of(null,
record3));
PCollection<Row> input = p.apply(Create.of(ROWS));
Schema errorSchema = ErrorHandling.errorSchema(BEAMSCHEMA);
@@ -238,7 +247,7 @@ public class KafkaWriteSchemaTransformProviderTest {
output.get(ERROR_TAG).setRowSchema(errorSchema);
output
.get(RECORD_OUTPUT_TAG)
- .setCoder(KvCoder.of(ByteArrayCoder.of(), AvroCoder.of(avroSchema)));
+ .setCoder(KvCoder.of(NullableCoder.of(ByteArrayCoder.of()),
AvroCoder.of(avroSchema)));
PAssert.that(output.get(RECORD_OUTPUT_TAG)).containsInAnyOrder(msg);
p.run().waitUntilFinish();
}