Skip to content

[SPARK-59662] Check for isRepartitionBatch can crash if metadata keys are missing - #58929

Open
ivasilyev1 wants to merge 2 commits into
apache:masterfrom
ivasilyev1:missing-shuffle-partitions
Open

ivasilyev1 wants to merge 2 commits into
apache:masterfrom
ivasilyev1:missing-shuffle-partitions

Conversation

@ivasilyev1

Copy link
Copy Markdown

What changes were proposed in this pull request?

This PR changes 2 things:

  • In OffsetSeq emit the rebinded SQLConf for both OffsetSeqMetadata and OffsetSeqMetadataV2
  • In OfflineStateRepartitionUtils.isRepartitionBatch add matched case to ensure shufflePartitions and previousShufflePartitions exist

Why are the changes needed?

These changes are necessary because in the current implementation OfflineStateRepartitionUtils.isRepartitionBatch assumes that spark.sql.shuffle.partitions will always be present in the offset metadata and accesses the keys directly with .get.

With the changes in offset log format v2, it is possible for the offset log to not contain this key if the current batch did not successfully complete. When this happens, the checks in OfflineStateRepartitionUtils.isRepartitionBatch will cause spark to crash:

26/09/19 14:05:17 ERROR MicroBatchExecution: Query [id = 41da1d72-c034-4a36-9ae3-fdc0e1297405, runId = 4eb3a723-257b-45c7-a2e0-72dff79ebf5b] terminated with error
java.util.NoSuchElementException: None.get
        at scala.None$.get(Option.scala:627)
        at scala.None$.get(Option.scala:626)
        at org.apache.spark.sql.execution.streaming.state.OfflineStateRepartitionUtils$.isRepartitionBatch(OfflineStateRepartitionUtils.scala:54)
        at org.apache.spark.sql.execution.streaming.runtime.MicroBatchExecution.checkUnfinishedRepartitionBatch(MicroBatchExecution.scala:564)
        at org.apache.spark.sql.execution.streaming.runtime.MicroBatchExecution.initializeExecution(MicroBatchExecution.scala:494)
        at org.apache.spark.sql.execution.streaming.runtime.MicroBatchExecution.runActivatedStream(MicroBatchExecution.scala:604)
        at org.apache.spark.sql.execution.streaming.runtime.StreamExecution.$anonfun$runStream$1(StreamExecution.scala:353)
        at scala.runtime.java8.JFunction0$mcV$sp.apply(JFunction0$mcV$sp.scala:18)
        at org.apache.spark.sql.SparkSession.withActive(SparkSession.scala:810)
        at org.apache.spark.sql.execution.streaming.runtime.StreamExecution.org$apache$spark$sql$execution$streaming$runtime$StreamExecution$$runStream(StreamExecution.scala:313)
        at org.apache.spark.sql.execution.streaming.runtime.StreamExecution$$anon$1.run(StreamExecution.scala:236)
Traceback (most recent call last): 

This behavior can be recreated with the following script to reliably trigger the crash:

import time
from pathlib import Path
from pyspark.sql import SparkSession


def get_spark():
    spark = (
        SparkSession.builder
        .master("local[2]")
        .appName("offset-log-bug")
        .config("spark.sql.streaming.offsetLog.formatVersion", "2")
        .config("spark.sql.streaming.checkUnfinishedRepartitionOnRestart", "true")
        .getOrCreate()
    )
    return spark

def fail_on_second_batch(batch_df, batch_id):
    if batch_id == 1:
        raise RuntimeError("Failing on second batch")
    batch_df.count()

def main():
    input_dir = Path("input")
    input_dir.mkdir(exist_ok=True)

    ( input_dir/"file1.txt" ).write_text("1\n")
    ( input_dir/"file2.txt" ).write_text("2\n")

    spark = get_spark()
    df = spark.readStream.format("text").option("maxFilesPerTrigger", "1").load("input")

    query = (
        df.writeStream
        .foreachBatch(fail_on_second_batch)
        .option("checkpointLocation", "checkpoint")
        .trigger(availableNow=True)
        .start()
    )

    try:
        query.awaitTermination()
    except:
        print("#####################")
        print("Expected first run failure")
        print("Shutting down spark and recreating")
        print("Continue in 5 seconds")
        print("#####################")
        time.sleep(5)
    finally:
        query.stop()
        spark.stop()


    spark = get_spark()
    df = (spark.readStream.format("text").option("maxFilesPerTrigger", "1").load("input"))

    query = (
        df.writeStream
        .foreachBatch(lambda batch_df, batch_id: batch_df.count())
        .option("checkpointLocation", "checkpoint")
        .trigger(availableNow=True)
        .start()
    )
    query.awaitTermination()


if __name__ == "__main__":
    main()

Does this PR introduce any user-facing change?

I am unsure if this counts as a user-facing change, but the noticeable change would be that spark.sql.shuffle.partitions key is always present in the offset metadata

How was this patch tested?

I tested the crash condition with the above python script. After introducing these change and compiling a spark distribution, the query no longer crashes and instead finishes successfully

Was this patch authored or co-authored using generative AI tooling?

Generated-by: GPT 5.6-Luna

@ivasilyev1

Copy link
Copy Markdown
Author

Please let me know if any changes are needed to conform with existing patterns as this is my first contribution to spark.

Technically the changed files both independently prevent the crash from happening, but I figured that both changes made sense to include.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant