Skip to content

Commit 0b407cd

Browse files
authored
Fix lost state updates in Set state (#39175)
* Add a test to reproduce the problem. * Fix race condition in SetState compaction by awaiting outstanding state requests * Reformat * Simply logic * Fix lints * Do not force to use prism for the new test.
1 parent 2940f8f commit 0b407cd

2 files changed

Lines changed: 95 additions & 11 deletions

File tree

sdks/python/apache_beam/runners/worker/bundle_processor.py

Lines changed: 26 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -641,6 +641,8 @@ def __init__(
641641
self._value_coder = value_coder
642642
self._cleared = False
643643
self._added_elements: set[Any] = set()
644+
# Track outstanding async state requests to await them at commit time.
645+
self._futures = []
644646

645647
def _compact_data(self, rewrite=True):
646648
accumulator = set(
@@ -650,9 +652,11 @@ def _compact_data(self, rewrite=True):
650652
self._added_elements))
651653

652654
if rewrite and accumulator:
653-
self._state_handler.clear(self._state_key)
654-
self._state_handler.extend(
655-
self._state_key, self._value_coder.get_impl(), accumulator)
655+
# Compaction writes are asynchronous; queue them so they are not lost.
656+
self._futures.append(self._state_handler.clear(self._state_key))
657+
self._futures.append(
658+
self._state_handler.extend(
659+
self._state_key, self._value_coder.get_impl(), accumulator))
656660

657661
# Since everthing is already committed so we can safely reinitialize
658662
# added_elements here.
@@ -666,7 +670,7 @@ def read(self) -> set[Any]:
666670
def add(self, value: Any) -> None:
667671
if self._cleared:
668672
# This is a good time explicitly clear.
669-
self._state_handler.clear(self._state_key)
673+
self._futures.append(self._state_handler.clear(self._state_key))
670674
self._cleared = False
671675

672676
self._added_elements.add(value)
@@ -678,15 +682,26 @@ def clear(self) -> None:
678682
self._added_elements = set()
679683

680684
def commit(self) -> None:
681-
to_await = None
682685
if self._cleared:
683-
to_await = self._state_handler.clear(self._state_key)
686+
self._futures.append(self._state_handler.clear(self._state_key))
687+
self._cleared = False
684688
if self._added_elements:
685-
to_await = self._state_handler.extend(
686-
self._state_key, self._value_coder.get_impl(), self._added_elements)
687-
if to_await:
688-
# To commit, we need to wait on the last state request future to complete.
689-
to_await.get()
689+
self._futures.append(
690+
self._state_handler.extend(
691+
self._state_key,
692+
self._value_coder.get_impl(),
693+
self._added_elements))
694+
self._added_elements = set()
695+
696+
# Block on all outstanding async state requests to ensure data is committed.
697+
# We must swap and clear self._futures before awaiting them. Awaiting a future
698+
# yields control, during which new futures could be appended to self._futures.
699+
all_futures = self._futures
700+
self._futures = []
701+
702+
for f in all_futures:
703+
if f:
704+
f.get()
690705

691706

692707
class RangeSet:

sdks/python/apache_beam/transforms/userstate_test.py

Lines changed: 69 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,9 @@
1818
"""Unit tests for the Beam State and Timer API interfaces."""
1919
# pytype: skip-file
2020

21+
import queue
22+
import threading
23+
import time
2124
import unittest
2225
from typing import Any
2326

@@ -34,6 +37,8 @@
3437
from apache_beam.portability.api import beam_runner_api_pb2
3538
from apache_beam.runners import pipeline_context
3639
from apache_beam.runners.common import DoFnSignature
40+
from apache_beam.runners.worker.sdk_worker import GrpcStateHandler
41+
from apache_beam.runners.worker.sdk_worker import _Future
3742
from apache_beam.testing.test_pipeline import TestPipeline
3843
from apache_beam.testing.test_stream import TestStream
3944
from apache_beam.testing.util import assert_that
@@ -719,6 +724,70 @@ def process(self, element, set_state=beam.DoFn.StateParam(SET_STATE)):
719724
actual_values = (values | beam.ParDo(SetStatefulDoFn()))
720725
assert_that(actual_values, equal_to([1, 3, 6, 10, 10]))
721726

727+
# Mock random to always return 1.0 to force compaction on every add.
728+
@mock.patch(
729+
'apache_beam.runners.worker.bundle_processor.random.random', lambda: 1.0)
730+
def test_stateful_set_state_compaction_race_portably(self):
731+
old_request = GrpcStateHandler._request
732+
request_queue = queue.Queue()
733+
734+
def worker():
735+
while True:
736+
handler, request, future, instruction_id = request_queue.get()
737+
time.sleep(0.1) # Simulate latency for each request sequentially.
738+
handler._context.process_instruction_id = instruction_id
739+
underlying_future = old_request(handler, request)
740+
underlying_future.wait()
741+
future.set(underlying_future.get())
742+
request_queue.task_done()
743+
744+
t = threading.Thread(target=worker, daemon=True)
745+
t.start()
746+
747+
def delayed_request(self, request):
748+
if request.HasField('append') or request.HasField('clear'):
749+
future = _Future()
750+
instruction_id = getattr(self._context, 'process_instruction_id', None)
751+
request_queue.put((self, request, future, instruction_id))
752+
return future
753+
else:
754+
return old_request(self, request)
755+
756+
GrpcStateHandler._request = delayed_request
757+
758+
class SetStatefulDoFn(beam.DoFn):
759+
760+
SET_STATE = SetStateSpec('buffer', VarIntCoder())
761+
762+
def process(self, element, set_state=beam.DoFn.StateParam(SET_STATE)):
763+
_, value = element
764+
aggregated_value = 0
765+
set_state.add(value)
766+
for saved_value in set_state.read():
767+
aggregated_value += saved_value
768+
yield aggregated_value
769+
770+
try:
771+
options = PipelineOptions([
772+
'--max_cache_memory_usage_mb=100',
773+
'--environment_type=LOOPBACK',
774+
])
775+
with TestPipeline(options=options) as p:
776+
test_stream = (
777+
TestStream(
778+
coder=beam.coders.TupleCoder((
779+
beam.coders.StrUtf8Coder(), beam.coders.VarIntCoder()
780+
))).advance_watermark_to(10).add_elements([
781+
('key', 1)
782+
]).advance_watermark_to(20).add_elements([
783+
('key', 2)
784+
]).advance_watermark_to(30))
785+
actual_values = (p | test_stream | beam.ParDo(SetStatefulDoFn()))
786+
assert_that(actual_values, equal_to([1, 3]))
787+
788+
finally:
789+
GrpcStateHandler._request = old_request
790+
722791
def test_stateful_set_state_clean_portably(self):
723792
class SetStateClearingStatefulDoFn(beam.DoFn):
724793

0 commit comments

Comments
 (0)