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();
   }

Reply via email to