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 e2c9256725c [KafkaIO] Add `with_gcp_adc` and `num_partitions` options 
to Kafka SchemaTransforms and YAML (#40181)
e2c9256725c is described below

commit e2c9256725c84e2ae1f5b616ef440252b370dd51
Author: Yi Hu <[email protected]>
AuthorDate: Mon Oct 5 15:25:23 2026 -0400

    [KafkaIO] Add `with_gcp_adc` and `num_partitions` options to Kafka 
SchemaTransforms and YAML (#40181)
    
    * [KafkaIO] Add `with_gcp_adc` and `num_partitions` options to Kafka 
SchemaTransforms and Beam YAML
    
    - Expose `with_gcp_adc` (`withGcpAdc`) in KafkaIO read/write 
SchemaTransform, Python xlang IO and YAML
    
    - Add `num_partitions` in KafkaIO read SchemaTransform, Python xlang IO and 
YAML. This avoids the need to Kafka instance connection at pipeline submission 
time for Dataflow runner v1 (Streaming runner) if using SchemaTransform
---
 .../KafkaReadSchemaTransformConfiguration.java     | 19 +++++++
 .../io/kafka/KafkaReadSchemaTransformProvider.java | 28 ++++++++--
 .../kafka/KafkaWriteSchemaTransformProvider.java   | 52 ++++++++++++------
 .../KafkaReadSchemaTransformProviderTest.java      | 64 +++++++++++++++++++++-
 .../KafkaWriteSchemaTransformProviderTest.java     | 13 ++++-
 sdks/python/apache_beam/yaml/standard_io.yaml      |  3 +
 6 files changed, 152 insertions(+), 27 deletions(-)

diff --git 
a/sdks/java/io/kafka/src/main/java/org/apache/beam/sdk/io/kafka/KafkaReadSchemaTransformConfiguration.java
 
b/sdks/java/io/kafka/src/main/java/org/apache/beam/sdk/io/kafka/KafkaReadSchemaTransformConfiguration.java
index 0cf40f9b7eb..aad7d742769 100644
--- 
a/sdks/java/io/kafka/src/main/java/org/apache/beam/sdk/io/kafka/KafkaReadSchemaTransformConfiguration.java
+++ 
b/sdks/java/io/kafka/src/main/java/org/apache/beam/sdk/io/kafka/KafkaReadSchemaTransformConfiguration.java
@@ -198,6 +198,21 @@ public abstract class 
KafkaReadSchemaTransformConfiguration {
   @Nullable
   public abstract Boolean getRedistributeByRecordKey();
 
+  @SchemaFieldDescription(
+      "Whether to use Google Application Default Credentials (ADC) for 
authenticating with a "
+          + "Google Managed Kafka cluster.")
+  @SchemaFieldNumber("17")
+  @Nullable
+  public abstract Boolean getWithGcpAdc();
+
+  @SchemaFieldDescription(
+      "The number of partitions to read from the Kafka topic. If specified, 
partitions"
+          + " 0 to numPartitions-1 will be assigned statically without 
querying Kafka at pipeline"
+          + " construction time.")
+  @SchemaFieldNumber("18")
+  @Nullable
+  public abstract Integer getNumPartitions();
+
   /** Builder for the {@link KafkaReadSchemaTransformConfiguration}. */
   @AutoValue.Builder
   public abstract static class Builder {
@@ -238,6 +253,10 @@ public abstract class 
KafkaReadSchemaTransformConfiguration {
 
     public abstract Builder setRedistributeByRecordKey(Boolean 
redistributeByRecordKey);
 
+    public abstract Builder setWithGcpAdc(Boolean withGcpAdc);
+
+    public abstract Builder setNumPartitions(Integer numPartitions);
+
     /** Builds a {@link KafkaReadSchemaTransformConfiguration} instance. */
     public abstract KafkaReadSchemaTransformConfiguration build();
   }
diff --git 
a/sdks/java/io/kafka/src/main/java/org/apache/beam/sdk/io/kafka/KafkaReadSchemaTransformProvider.java
 
b/sdks/java/io/kafka/src/main/java/org/apache/beam/sdk/io/kafka/KafkaReadSchemaTransformProvider.java
index c5764b39bc6..aa8b627c63e 100644
--- 
a/sdks/java/io/kafka/src/main/java/org/apache/beam/sdk/io/kafka/KafkaReadSchemaTransformProvider.java
+++ 
b/sdks/java/io/kafka/src/main/java/org/apache/beam/sdk/io/kafka/KafkaReadSchemaTransformProvider.java
@@ -31,6 +31,7 @@ import java.nio.channels.WritableByteChannel;
 import java.nio.charset.StandardCharsets;
 import java.nio.file.Files;
 import java.nio.file.Path;
+import java.util.ArrayList;
 import java.util.Arrays;
 import java.util.HashMap;
 import java.util.List;
@@ -69,6 +70,7 @@ import 
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.Lists;
 import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.Maps;
 import org.apache.kafka.clients.consumer.Consumer;
 import org.apache.kafka.clients.consumer.ConsumerConfig;
+import org.apache.kafka.common.TopicPartition;
 import org.apache.kafka.common.serialization.ByteArrayDeserializer;
 import org.joda.time.Duration;
 import org.slf4j.Logger;
@@ -166,6 +168,20 @@ public class KafkaReadSchemaTransformProvider
       return SchemaRegistryProvider.UNSPECIFIED;
     }
 
+    private static <K, V> KafkaIO.Read<K, V> applyTopicOrPartitions(
+        KafkaIO.Read<K, V> kafkaRead, KafkaReadSchemaTransformConfiguration 
configuration) {
+      Integer numPartitions = configuration.getNumPartitions();
+      String topic = configuration.getTopic();
+      if (numPartitions != null && numPartitions > 0) {
+        List<TopicPartition> topicPartitions = new ArrayList<>(numPartitions);
+        for (int i = 0; i < numPartitions; i++) {
+          topicPartitions.add(new TopicPartition(topic, i));
+        }
+        return kafkaRead.withTopicPartitions(topicPartitions);
+      }
+      return kafkaRead.withTopic(topic);
+    }
+
     private static <K, V> KafkaIO.Read<K, V> applyRedistributeSettings(
         KafkaIO.Read<K, V> kafkaRead, KafkaReadSchemaTransformConfiguration 
configuration) {
       Boolean redistribute = configuration.getRedistributed();
@@ -221,12 +237,14 @@ public class KafkaReadSchemaTransformProvider
         KafkaIO.Read<byte[], GenericRecord> kafkaRead;
 
         kafkaRead =
-            KafkaIO.<byte[], GenericRecord>read()
-                .withTopic(configuration.getTopic())
+            applyTopicOrPartitions(KafkaIO.<byte[], GenericRecord>read(), 
configuration)
                 .withConsumerFactoryFn(new ConsumerFactoryWithGcsTrustStores())
                 .withBootstrapServers(configuration.getBootstrapServers())
                 .withConsumerConfigUpdates(consumerConfigs)
                 .withKeyDeserializer(ByteArrayDeserializer.class);
+        if (Boolean.TRUE.equals(configuration.getWithGcpAdc())) {
+          kafkaRead = kafkaRead.withGCPApplicationDefaultCredentials();
+        }
 
         SchemaRegistryProvider provider = 
getSchemaRegistryProvider(confluentSchemaRegUrl);
         switch (provider) {
@@ -300,11 +318,13 @@ public class KafkaReadSchemaTransformProvider
       }
 
       KafkaIO.Read<byte[], byte[]> kafkaRead =
-          KafkaIO.readBytes()
+          applyTopicOrPartitions(KafkaIO.readBytes(), configuration)
               .withConsumerConfigUpdates(consumerConfigs)
               .withConsumerFactoryFn(new ConsumerFactoryWithGcsTrustStores())
-              .withTopic(configuration.getTopic())
               .withBootstrapServers(configuration.getBootstrapServers());
+      if (Boolean.TRUE.equals(configuration.getWithGcpAdc())) {
+        kafkaRead = kafkaRead.withGCPApplicationDefaultCredentials();
+      }
       Integer maxReadTimeSeconds = configuration.getMaxReadTimeSeconds();
       if (maxReadTimeSeconds != null) {
         kafkaRead = 
kafkaRead.withMaxReadTime(Duration.standardSeconds(maxReadTimeSeconds));
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 e59159ba5b8..6dc45c8253d 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
@@ -256,17 +256,20 @@ public class KafkaWriteSchemaTransformProvider
                                 handleErrors))
                         .withOutputTags(RECORD_OUTPUT_TAG, 
TupleTagList.of(ERROR_TAG)));
         HashMap<String, Object> producerConfig = new 
HashMap<>(configOverrides);
+        KafkaIO.Write<byte[], GenericRecord> kafkaWrite =
+            KafkaIO.<byte[], GenericRecord>write()
+                .withTopic(configuration.getTopic())
+                .withBootstrapServers(configuration.getBootstrapServers())
+                .withProducerConfigUpdates(producerConfig)
+                .withKeySerializer(ByteArraySerializer.class)
+                .withValueSerializer((Class) KafkaAvroSerializer.class);
+        if (Boolean.TRUE.equals(configuration.getWithGcpAdc())) {
+          kafkaWrite = kafkaWrite.withGCPApplicationDefaultCredentials();
+        }
         outputTuple
             .get(RECORD_OUTPUT_TAG)
             .setCoder(KvCoder.of(NullableCoder.of(ByteArrayCoder.of()), 
AvroCoder.of(avroSchema)))
-            .apply(
-                "Map Rows to GenericRecords",
-                KafkaIO.<byte[], GenericRecord>write()
-                    .withTopic(configuration.getTopic())
-                    .withBootstrapServers(configuration.getBootstrapServers())
-                    .withProducerConfigUpdates(producerConfig)
-                    .withKeySerializer(ByteArraySerializer.class)
-                    .withValueSerializer((Class) KafkaAvroSerializer.class));
+            .apply("Map Rows to GenericRecords", kafkaWrite);
       } else {
         outputTuple =
             input
@@ -278,19 +281,23 @@ public class KafkaWriteSchemaTransformProvider
                                 "Kafka-write-error-counter", toBytesFn, 
errorSchema, handleErrors))
                         .withOutputTags(OUTPUT_TAG, 
TupleTagList.of(ERROR_TAG)));
 
+        KafkaIO.Write<byte[], byte[]> kafkaWrite =
+            KafkaIO.<byte[], byte[]>write()
+                .withTopic(configuration.getTopic())
+                .withBootstrapServers(configuration.getBootstrapServers())
+                .withProducerConfigUpdates(
+                    configOverrides == null
+                        ? new HashMap<>()
+                        : new HashMap<String, Object>(configOverrides))
+                .withKeySerializer(ByteArraySerializer.class)
+                .withValueSerializer(ByteArraySerializer.class);
+        if (Boolean.TRUE.equals(configuration.getWithGcpAdc())) {
+          kafkaWrite = kafkaWrite.withGCPApplicationDefaultCredentials();
+        }
         outputTuple
             .get(OUTPUT_TAG)
             .setCoder(KvCoder.of(NullableCoder.of(ByteArrayCoder.of()), 
ByteArrayCoder.of()))
-            .apply(
-                KafkaIO.<byte[], byte[]>write()
-                    .withTopic(configuration.getTopic())
-                    .withBootstrapServers(configuration.getBootstrapServers())
-                    .withProducerConfigUpdates(
-                        configOverrides == null
-                            ? new HashMap<>()
-                            : new HashMap<String, Object>(configOverrides))
-                    .withKeySerializer(ByteArraySerializer.class)
-                    .withValueSerializer(ByteArraySerializer.class));
+            .apply(kafkaWrite);
       }
 
       // TODO: include output from KafkaIO Write once updated from PDone
@@ -380,6 +387,13 @@ public class KafkaWriteSchemaTransformProvider
     @Nullable
     public abstract String getSchema();
 
+    @SchemaFieldDescription(
+        "Whether to use Google Application Default Credentials (ADC) for 
authenticating with a "
+            + "Google Managed Kafka cluster.")
+    @SchemaFieldNumber("8")
+    @Nullable
+    public abstract Boolean getWithGcpAdc();
+
     public static Builder builder() {
       return new 
AutoValue_KafkaWriteSchemaTransformProvider_KafkaWriteSchemaTransformConfiguration
           .Builder();
@@ -403,6 +417,8 @@ public class KafkaWriteSchemaTransformProvider
 
       public abstract Builder setSchema(String schema);
 
+      public abstract Builder setWithGcpAdc(Boolean withGcpAdc);
+
       public abstract KafkaWriteSchemaTransformConfiguration build();
     }
   }
diff --git 
a/sdks/java/io/kafka/src/test/java/org/apache/beam/sdk/io/kafka/KafkaReadSchemaTransformProviderTest.java
 
b/sdks/java/io/kafka/src/test/java/org/apache/beam/sdk/io/kafka/KafkaReadSchemaTransformProviderTest.java
index 9d276fa0e55..d5a0b7b152b 100644
--- 
a/sdks/java/io/kafka/src/test/java/org/apache/beam/sdk/io/kafka/KafkaReadSchemaTransformProviderTest.java
+++ 
b/sdks/java/io/kafka/src/test/java/org/apache/beam/sdk/io/kafka/KafkaReadSchemaTransformProviderTest.java
@@ -18,7 +18,9 @@
 package org.apache.beam.sdk.io.kafka;
 
 import static org.junit.Assert.assertEquals;
+import static org.junit.Assert.assertNotNull;
 import static org.junit.Assert.assertThrows;
+import static org.junit.Assert.assertTrue;
 
 import java.io.IOException;
 import java.nio.charset.StandardCharsets;
@@ -26,10 +28,12 @@ import java.util.Arrays;
 import java.util.List;
 import java.util.Objects;
 import java.util.ServiceLoader;
+import java.util.concurrent.atomic.AtomicReference;
 import java.util.stream.Collectors;
 import java.util.stream.StreamSupport;
 import org.apache.beam.sdk.Pipeline;
 import org.apache.beam.sdk.managed.Managed;
+import org.apache.beam.sdk.runners.TransformHierarchy;
 import org.apache.beam.sdk.schemas.NoSuchSchemaException;
 import org.apache.beam.sdk.schemas.Schema;
 import org.apache.beam.sdk.schemas.SchemaRegistry;
@@ -41,6 +45,7 @@ import org.apache.beam.sdk.values.PCollectionRowTuple;
 import 
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.Lists;
 import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.Sets;
 import 
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.io.ByteStreams;
+import org.apache.kafka.common.TopicPartition;
 import org.junit.Test;
 import org.junit.runner.RunWith;
 import org.junit.runners.JUnit4;
@@ -138,7 +143,9 @@ public class KafkaReadSchemaTransformProviderTest {
             "allow_duplicates",
             "offset_deduplication",
             "redistribute_num_keys",
-            "redistribute_by_record_key"),
+            "redistribute_by_record_key",
+            "with_gcp_adc",
+            "num_partitions"),
         kafkaProvider.configurationSchema().getFields().stream()
             .map(field -> field.getName())
             .collect(Collectors.toSet()));
@@ -361,7 +368,12 @@ public class KafkaReadSchemaTransformProviderTest {
                 + "schema: '"
                 + PROTO_SCHEMA
                 + "'\n"
-                + "message_name: MyMessage");
+                + "message_name: MyMessage",
+            "topic: topic_6\n"
+                + "bootstrap_servers: some bootstrap\n"
+                + "format: RAW\n"
+                + "with_gcp_adc: true\n"
+                + "num_partitions: 3");
 
     for (String config : configs) {
       // Kafka Read SchemaTransform gets built in 
ManagedSchemaTransformProvider's expand
@@ -371,6 +383,42 @@ public class KafkaReadSchemaTransformProviderTest {
     }
   }
 
+  @Test
+  public void testBuildTransformWithNumPartitions() {
+    KafkaReadSchemaTransformProvider kafkaProvider = new 
KafkaReadSchemaTransformProvider();
+    SchemaTransform transformWithPartitions =
+        kafkaProvider.from(
+            KafkaReadSchemaTransformConfiguration.builder()
+                .setTopic("anytopic")
+                .setBootstrapServers("anybootstrap")
+                .setFormat("RAW")
+                .setNumPartitions(3)
+                .build());
+    Pipeline pipelineWithPartitions = Pipeline.create();
+    
transformWithPartitions.expand(PCollectionRowTuple.empty(pipelineWithPartitions));
+    AtomicReference<KafkaIO.Read<?, ?>> readWithPartitions = new 
AtomicReference<>();
+    pipelineWithPartitions.traverseTopologically(
+        new Pipeline.PipelineVisitor.Defaults() {
+          @Override
+          public CompositeBehavior 
enterCompositeTransform(TransformHierarchy.Node node) {
+            if (node.getTransform() instanceof KafkaIO.Read) {
+              readWithPartitions.set((KafkaIO.Read<?, ?>) node.getTransform());
+            }
+            return CompositeBehavior.ENTER_TRANSFORM;
+          }
+        });
+    assertNotNull(readWithPartitions.get());
+    assertEquals(
+        Arrays.asList(
+            new TopicPartition("anytopic", 0),
+            new TopicPartition("anytopic", 1),
+            new TopicPartition("anytopic", 2)),
+        readWithPartitions.get().getTopicPartitions());
+    assertTrue(
+        readWithPartitions.get().getTopics() == null
+            || readWithPartitions.get().getTopics().isEmpty());
+  }
+
   // This test verifies that the schema for 
KafkaReadSchemaTransformConfiguration is correctly
   // generated. This schema is used when KafkaReadSchemaTransformConfiguration 
are
   // serialized/deserialized with
@@ -380,7 +428,7 @@ public class KafkaReadSchemaTransformProviderTest {
     Schema schema =
         
SchemaRegistry.createDefault().getSchema(KafkaReadSchemaTransformConfiguration.class);
 
-    assertEquals(17, schema.getFieldCount());
+    assertEquals(19, schema.getFieldCount());
 
     // Check field name, type, and nullability. Descriptions are not checked 
as they are not
     // critical for serialization.
@@ -478,5 +526,15 @@ public class KafkaReadSchemaTransformProviderTest {
         Schema.Field.nullable("redistributeByRecordKey", 
Schema.FieldType.BOOLEAN)
             .withDescription(schema.getField(16).getDescription()),
         schema.getField(16));
+
+    assertEquals(
+        Schema.Field.nullable("withGcpAdc", Schema.FieldType.BOOLEAN)
+            .withDescription(schema.getField(17).getDescription()),
+        schema.getField(17));
+
+    assertEquals(
+        Schema.Field.nullable("numPartitions", Schema.FieldType.INT32)
+            .withDescription(schema.getField(18).getDescription()),
+        schema.getField(18));
   }
 }
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 ef53ff0bb83..f76c26af946 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
@@ -269,7 +269,11 @@ public class KafkaWriteSchemaTransformProviderTest {
                 + "schema: '"
                 + PROTO_SCHEMA
                 + "'\n"
-                + "message_name: MyMessage");
+                + "message_name: MyMessage",
+            "topic: topic_4\n"
+                + "bootstrap_servers: some bootstrap\n"
+                + "format: RAW\n"
+                + "with_gcp_adc: true");
 
     for (String config : configs) {
       // Kafka Write SchemaTransform gets built in 
ManagedSchemaTransformProvider's expand
@@ -313,7 +317,7 @@ public class KafkaWriteSchemaTransformProviderTest {
 
     System.out.println("schema = " + schema);
 
-    assertEquals(8, schema.getFieldCount());
+    assertEquals(9, schema.getFieldCount());
 
     // Check field name, type, and nullability. Descriptions are not checked 
as they are not
     // critical for serialization.
@@ -365,5 +369,10 @@ public class KafkaWriteSchemaTransformProviderTest {
         Schema.Field.nullable("schema", Schema.FieldType.STRING)
             .withDescription(schema.getField(7).getDescription()),
         schema.getField(7));
+
+    assertEquals(
+        Schema.Field.nullable("withGcpAdc", Schema.FieldType.BOOLEAN)
+            .withDescription(schema.getField(8).getDescription()),
+        schema.getField(8));
   }
 }
diff --git a/sdks/python/apache_beam/yaml/standard_io.yaml 
b/sdks/python/apache_beam/yaml/standard_io.yaml
index 825cd204ed6..15aa9186c0f 100644
--- a/sdks/python/apache_beam/yaml/standard_io.yaml
+++ b/sdks/python/apache_beam/yaml/standard_io.yaml
@@ -110,6 +110,8 @@
         'file_descriptor_path': 'file_descriptor_path'
         'message_name': 'message_name'
         'max_read_time_seconds': 'max_read_time_seconds'
+        'with_gcp_adc': 'with_gcp_adc'
+        'num_partitions': 'num_partitions'
       'WriteToKafka':
         'format': 'format'
         'topic': 'topic'
@@ -119,6 +121,7 @@
         'file_descriptor_path': 'file_descriptor_path'
         'message_name': 'message_name'
         'schema': 'schema'
+        'with_gcp_adc': 'with_gcp_adc'
     underlying_provider:
       type: beamJar
       transforms:

Reply via email to