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() {

Reply via email to