Skip to content

Commit fb95a66

Browse files
tkaymakclaude
andcommitted
refactor: make shared Spark source compatible with Scala 2.12 and 2.13
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
1 parent 4d22554 commit fb95a66

7 files changed

Lines changed: 27 additions & 16 deletions

File tree

runners/spark/src/main/java/org/apache/beam/runners/spark/coders/SparkRunnerKryoRegistrator.java

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,6 @@
3030
import org.apache.beam.sdk.values.TupleTag;
3131
import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.HashBasedTable;
3232
import org.apache.spark.serializer.KryoRegistrator;
33-
import scala.collection.mutable.WrappedArray;
3433

3534
/**
3635
* Custom {@link KryoRegistrator}s for Beam's Spark runner needs and registering used class in spark
@@ -61,7 +60,16 @@ public void registerClasses(Kryo kryo) {
6160
kryo.register(PaneInfo.class);
6261
kryo.register(StateAndTimers.class);
6362
kryo.register(TupleTag.class);
64-
kryo.register(WrappedArray.ofRef.class);
63+
// Scala 2.12 uses WrappedArray$ofRef, Scala 2.13 renamed it to ArraySeq$ofRef
64+
try {
65+
kryo.register(Class.forName("scala.collection.mutable.ArraySeq$ofRef"));
66+
} catch (ClassNotFoundException e) {
67+
try {
68+
kryo.register(Class.forName("scala.collection.mutable.WrappedArray$ofRef"));
69+
} catch (ClassNotFoundException ignored) {
70+
// Neither class found; skip registration
71+
}
72+
}
6573

6674
try {
6775
kryo.register(

runners/spark/src/main/java/org/apache/beam/runners/spark/io/SourceRDD.java

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -50,7 +50,7 @@
5050
import org.slf4j.Logger;
5151
import org.slf4j.LoggerFactory;
5252
import scala.Option;
53-
import scala.collection.JavaConversions;
53+
import scala.collection.JavaConverters;
5454

5555
/** Classes implementing Beam {@link Source} {@link RDD}s. */
5656
@SuppressWarnings({
@@ -75,7 +75,7 @@ public static class Bounded<T> extends RDD<WindowedValue<T>> {
7575

7676
// to satisfy Scala API.
7777
private static final scala.collection.immutable.Seq<Dependency<?>> NIL =
78-
JavaConversions.asScalaBuffer(Collections.<Dependency<?>>emptyList()).toList();
78+
JavaConverters.asScalaBuffer(Collections.<Dependency<?>>emptyList()).toList();
7979

8080
public Bounded(
8181
SparkContext sc,
@@ -148,7 +148,7 @@ public scala.collection.Iterator<WindowedValue<T>> compute(
148148
final Iterator<WindowedValue<T>> readerIterator =
149149
new ReaderToIteratorAdapter<>(metricsContainer, reader);
150150

151-
return new InterruptibleIterator<>(context, JavaConversions.asScalaIterator(readerIterator));
151+
return new InterruptibleIterator<>(context, JavaConverters.asScalaIterator(readerIterator));
152152
}
153153

154154
/**
@@ -299,7 +299,7 @@ public static class Unbounded<T, CheckpointMarkT extends UnboundedSource.Checkpo
299299

300300
// to satisfy Scala API.
301301
private static final scala.collection.immutable.List<Dependency<?>> NIL =
302-
JavaConversions.asScalaBuffer(Collections.<Dependency<?>>emptyList()).toList();
302+
JavaConverters.asScalaBuffer(Collections.<Dependency<?>>emptyList()).toList();
303303

304304
public Unbounded(
305305
SparkContext sc,
@@ -344,7 +344,7 @@ public scala.collection.Iterator<scala.Tuple2<Source<T>, CheckpointMarkT>> compu
344344
(CheckpointableSourcePartition<T, CheckpointMarkT>) split;
345345
scala.Tuple2<Source<T>, CheckpointMarkT> tuple2 =
346346
new scala.Tuple2<>(partition.getSource(), partition.checkpointMark);
347-
return JavaConversions.asScalaIterator(Collections.singleton(tuple2).iterator());
347+
return JavaConverters.asScalaIterator(Collections.singleton(tuple2).iterator());
348348
}
349349
}
350350

runners/spark/src/main/java/org/apache/beam/runners/spark/io/SparkUnboundedSource.java

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -186,7 +186,7 @@ public Duration slideDuration() {
186186

187187
@Override
188188
public scala.collection.immutable.List<DStream<?>> dependencies() {
189-
return scala.collection.JavaConversions.asScalaBuffer(
189+
return scala.collection.JavaConverters.asScalaBuffer(
190190
Collections.<DStream<?>>singletonList(parent))
191191
.toList();
192192
}

runners/spark/src/main/java/org/apache/beam/runners/spark/stateful/SparkGroupAlsoByWindowViaWindowSet.java

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -73,7 +73,7 @@
7373
import scala.Tuple2;
7474
import scala.Tuple3;
7575
import scala.collection.Iterator;
76-
import scala.collection.JavaConversions;
76+
import scala.collection.JavaConverters;
7777
import scala.collection.Seq;
7878
import scala.runtime.AbstractFunction1;
7979

@@ -238,7 +238,7 @@ private Collection<TimerInternals.TimerData> filterTimersEligibleForProcessing(
238238
// new input for key.
239239
try {
240240
final Iterable<WindowedValue<InputT>> elements =
241-
FluentIterable.from(JavaConversions.asJavaIterable(encodedElements))
241+
FluentIterable.from(JavaConverters.asJavaIterable(encodedElements))
242242
.transform(bytes -> CoderHelpers.fromByteArray(bytes, wvCoder));
243243

244244
LOG.trace("{}: input elements: {}", logPrefix, elements);
@@ -410,7 +410,7 @@ private Collection<TimerInternals.TimerData> filterTimersEligibleForProcessing(
410410
droppedDueToClosedWindow.inc(-droppedDueToClosedWindow.getCumulative());
411411
}
412412

413-
return scala.collection.JavaConversions.asScalaIterator(
413+
return JavaConverters.asScalaIterator(
414414
new UpdateStateByKeyOutputIterator(input, reduceFn, droppedDueToLateness));
415415
}
416416
}
@@ -522,7 +522,9 @@ JavaDStream<WindowedValue<KV<K, Iterable<InputT>>>> groupByKeyAndWindow(
522522
Tuple2</*K*/ ByteArray, Tuple2<StateAndTimers, /*WV<KV<K, Itr<I>>>*/ List<byte[]>>>>
523523
firedStream =
524524
pairDStream.updateStateByKey(
525-
updateFunc,
525+
// Raw cast to AbstractFunction1 suppresses Scala 2.12 (collection.Seq) vs
526+
// Scala 2.13 (immutable.Seq) type difference — safe at runtime due to erasure.
527+
(scala.runtime.AbstractFunction1) updateFunc,
526528
pairDStream.defaultPartitioner(pairDStream.defaultPartitioner$default$1()),
527529
true,
528530
JavaSparkContext$.MODULE$.fakeClassTag());

runners/spark/src/main/java/org/apache/beam/runners/spark/translation/SparkStreamingPortablePipelineTranslator.java

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -330,7 +330,8 @@ private static <T> void translateFlatten(
330330
}
331331
}
332332
// Unify streams into a single stream.
333-
unifiedStreams = context.getStreamingContext().union(JavaConverters.asScalaBuffer(dStreams));
333+
unifiedStreams =
334+
context.getStreamingContext().union(JavaConverters.asScalaBuffer(dStreams).toList());
334335
}
335336

336337
context.pushDataset(

runners/spark/src/main/java/org/apache/beam/runners/spark/translation/streaming/ParDoStateUpdateFn.java

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@
1919

2020
import java.io.Serializable;
2121
import java.util.Collection;
22+
import java.util.Collections;
2223
import java.util.Iterator;
2324
import java.util.List;
2425
import java.util.Map;
@@ -62,7 +63,6 @@
6263
import org.checkerframework.checker.nullness.qual.Nullable;
6364
import org.slf4j.Logger;
6465
import org.slf4j.LoggerFactory;
65-
import org.sparkproject.guava.collect.Iterators;
6666
import scala.Option;
6767
import scala.Tuple2;
6868
import scala.runtime.AbstractFunction3;
@@ -236,7 +236,7 @@ public TimerInternals timerInternals() {
236236
final byte[] byteValue = serializedValue.get();
237237
@Nullable WindowedValue<ValueT> windowedValue;
238238
@Nullable WindowedValue<KV<KeyT, ValueT>> keyedWindowedValue;
239-
Iterator<WindowedValue<KV<KeyT, ValueT>>> iterator = Iterators.emptyIterator();
239+
Iterator<WindowedValue<KV<KeyT, ValueT>>> iterator = Collections.emptyIterator();
240240
if (byteValue.length > 0) {
241241
windowedValue = CoderHelpers.fromByteArray(byteValue, this.wvCoder);
242242
keyedWindowedValue = windowedValue.withValue(KV.of(key, windowedValue.getValue()));

runners/spark/src/main/java/org/apache/beam/runners/spark/translation/streaming/StreamingTransformTranslator.java

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -306,7 +306,7 @@ public void evaluate(Flatten.PCollections<T> transform, EvaluationContext contex
306306
}
307307
// start by unifying streams into a single stream.
308308
JavaDStream<WindowedValue<T>> unifiedStreams =
309-
context.getStreamingContext().union(JavaConverters.asScalaBuffer(dStreams));
309+
context.getStreamingContext().union(JavaConverters.asScalaBuffer(dStreams).toList());
310310
context.putDataset(transform, new UnboundedDataset<>(unifiedStreams, streamingSources));
311311
}
312312

0 commit comments

Comments
 (0)