This is an automated email from the ASF dual-hosted git repository.
kennknowles 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 a15b880c53c Merge pull request #39043: Improve WithKeys coder
inference context
a15b880c53c is described below
commit a15b880c53c6d7f4bc6533ebce778881edb6e532
Author: ADITYA RAJ <[email protected]>
AuthorDate: Wed Jul 22 21:41:19 2026 +0530
Merge pull request #39043: Improve WithKeys coder inference context
---
.../org/apache/beam/sdk/transforms/WithKeys.java | 40 +++++++++++++++-------
.../apache/beam/sdk/transforms/WithKeysTest.java | 15 ++++++++
2 files changed, 42 insertions(+), 13 deletions(-)
diff --git
a/sdks/java/core/src/main/java/org/apache/beam/sdk/transforms/WithKeys.java
b/sdks/java/core/src/main/java/org/apache/beam/sdk/transforms/WithKeys.java
index 96072d8ec29..34af5811af1 100644
--- a/sdks/java/core/src/main/java/org/apache/beam/sdk/transforms/WithKeys.java
+++ b/sdks/java/core/src/main/java/org/apache/beam/sdk/transforms/WithKeys.java
@@ -110,27 +110,33 @@ public class WithKeys<K, V> extends
PTransform<PCollection<V>, PCollection<KV<K,
@Override
public PCollection<KV<K, V>> expand(PCollection<V> in) {
+ SerializableFunction<V, K> localFn = fn;
+ TypeDescriptor<V> inputType = in.getTypeDescriptor();
+ TypeDescriptor<KV<K, V>> outputType = getOutputTypeDescriptor(inputType);
PCollection<KV<K, V>> result =
in.apply(
"AddKeys",
- MapElements.via(
- new SimpleFunction<V, KV<K, V>>() {
- @Override
- public KV<K, V> apply(V element) {
- return KV.of(fn.apply(element), element);
- }
- }));
+ outputType == null
+ ? MapElements.via(
+ new SimpleFunction<V, KV<K, V>>() {
+ @Override
+ public KV<K, V> apply(V element) {
+ return KV.of(localFn.apply(element), element);
+ }
+ })
+ : MapElements.into(outputType)
+ .via(
+ (SerializableFunction<V, KV<K, V>>)
+ element -> KV.of(localFn.apply(element),
element)));
try {
- Coder<K> keyCoder;
CoderRegistry coderRegistry = in.getPipeline().getCoderRegistry();
- if (keyType == null) {
- keyCoder = coderRegistry.getOutputCoder(fn, in.getCoder());
+ if (outputType == null) {
+ Coder<K> keyCoder = coderRegistry.getOutputCoder(fn, in.getCoder());
+ result.setCoder(KvCoder.of(keyCoder, in.getCoder()));
} else {
- keyCoder = coderRegistry.getCoder(keyType);
+ result.setCoder(coderRegistry.getCoder(outputType,
checkNotNull(inputType), in.getCoder()));
}
- // TODO: Remove when we can set the coder inference context.
- result.setCoder(KvCoder.of(keyCoder, in.getCoder()));
} catch (CannotProvideCoderException exc) {
if (keyType != null) {
try {
@@ -151,4 +157,12 @@ public class WithKeys<K, V> extends
PTransform<PCollection<V>, PCollection<KV<K,
return result;
}
+
+ private @Nullable TypeDescriptor<KV<K, V>> getOutputTypeDescriptor(
+ @Nullable TypeDescriptor<V> inputType) {
+ if (keyType == null || inputType == null) {
+ return null;
+ }
+ return TypeDescriptors.kvs(keyType, inputType);
+ }
}
diff --git
a/sdks/java/core/src/test/java/org/apache/beam/sdk/transforms/WithKeysTest.java
b/sdks/java/core/src/test/java/org/apache/beam/sdk/transforms/WithKeysTest.java
index fd178f8e764..e16b0fb77c5 100644
---
a/sdks/java/core/src/test/java/org/apache/beam/sdk/transforms/WithKeysTest.java
+++
b/sdks/java/core/src/test/java/org/apache/beam/sdk/transforms/WithKeysTest.java
@@ -172,6 +172,21 @@ public class WithKeysTest {
p.run();
}
+ @Test
+ public void withKeyTypeShouldSetOutputTypeDescriptorFromInputType() {
+ PCollection<String> values =
+ p.apply(Create.of("1234", "3210").withType(TypeDescriptors.strings()));
+
+ PCollection<KV<Integer, String>> kvs =
+ values.apply(
+ WithKeys.of((SerializableFunction<String, Integer>)
Integer::valueOf)
+ .withKeyType(TypeDescriptors.integers()));
+
+ assertEquals(
+ TypeDescriptors.kvs(TypeDescriptors.integers(),
TypeDescriptors.strings()),
+ kvs.getTypeDescriptor());
+ }
+
@Test
@Category(NeedsRunner.class)
public void withLambdaAndNoTypeDescriptorShouldThrow() {