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: