From 60231dc8e57df6912275ceb9f9a0f046d35b6bf8 Mon Sep 17 00:00:00 2001 From: bvolpato Date: Thu, 3 Sep 2026 17:42:35 +0000 Subject: [PATCH 1/2] [Python] Fail bundles on state write errors --- CHANGES.md | 1 + .../runners/worker/bundle_processor.py | 35 +++--- .../runners/worker/bundle_processor_test.py | 112 ++++++++++++++++++ 3 files changed, 134 insertions(+), 14 deletions(-) diff --git a/CHANGES.md b/CHANGES.md index f0d5d06b9d9f..e5d0f727abc5 100644 --- a/CHANGES.md +++ b/CHANGES.md @@ -93,6 +93,7 @@ ## Bugfixes * (Java) Fixed the Spark runner firing processing-time timers in reverse timestamp order ([#39824](https://github.com/apache/beam/issues/39824)). +* (Python) State write failures reported by portable runners now fail the bundle instead of silently losing state updates. * (Python) Fixed incorrect profiler options handling on portable runners ([#39613](https://github.com/apache/beam/issues/39613)). * (Java) KafkaIO dynamic reads no longer require the obsolete `beam_fn_api` experiment ([#29998](https://github.com/apache/beam/issues/29998)). * (Prism) Self-checkpointing splittable DoFns now resume after their requested delay instead of immediately, so polling SDFs no longer busy-spin ([#39848](https://github.com/apache/beam/issues/39848)). diff --git a/sdks/python/apache_beam/runners/worker/bundle_processor.py b/sdks/python/apache_beam/runners/worker/bundle_processor.py index 851efc81221d..82b86e039c67 100644 --- a/sdks/python/apache_beam/runners/worker/bundle_processor.py +++ b/sdks/python/apache_beam/runners/worker/bundle_processor.py @@ -592,6 +592,17 @@ def __iter__(self) -> Iterator[Any]: coder_impl.FastPrimitivesCoderImpl.register_iterable_like_type(_ConcatIterable) +def _await_state_commit(futures) -> None: + # Resolve every write before reporting errors, including clears followed by + # appends. State futures return error responses rather than raising them. + responses = [future.get() for future in futures if future] + errors = [ + response.error for response in responses if response and response.error + ] + if errors: + raise RuntimeError('\n'.join(errors)) + + class SynchronousBagRuntimeState(userstate.BagRuntimeState): def __init__( self, @@ -619,15 +630,16 @@ def clear(self) -> None: self._added_elements = [] def commit(self) -> None: - to_await = None + futures = [] if self._cleared: - to_await = self._state_handler.clear(self._state_key) + futures.append(self._state_handler.clear(self._state_key)) if self._added_elements: - to_await = self._state_handler.extend( - self._state_key, self._value_coder.get_impl(), self._added_elements) - if to_await: - # To commit, we need to wait on the last state request future to complete. - to_await.get() + futures.append( + self._state_handler.extend( + self._state_key, + self._value_coder.get_impl(), + self._added_elements)) + _await_state_commit(futures) class SynchronousSetRuntimeState(userstate.SetRuntimeState): @@ -699,9 +711,7 @@ def commit(self) -> None: all_futures = self._futures self._futures = [] - for f in all_futures: - if f: - f.get() + _await_state_commit(all_futures) class RangeSet: @@ -870,10 +880,7 @@ def commit(self) -> None: self._state_key, self._elem_coder.get_impl(), items_to_add)) self._pending_adds = SortedDict() - if len(futures): - # To commit, we need to wait on every state request futures to complete. - for to_await in futures: - to_await.get() + _await_state_commit(futures) self._cleared = False diff --git a/sdks/python/apache_beam/runners/worker/bundle_processor_test.py b/sdks/python/apache_beam/runners/worker/bundle_processor_test.py index 0eb4dd9485fd..b0305e5e0609 100644 --- a/sdks/python/apache_beam/runners/worker/bundle_processor_test.py +++ b/sdks/python/apache_beam/runners/worker/bundle_processor_test.py @@ -20,6 +20,7 @@ import random import unittest +from unittest import mock import apache_beam as beam from apache_beam.coders import StrUtf8Coder @@ -30,6 +31,7 @@ from apache_beam.runners.portability.fn_api_runner.worker_handlers import StateServicer from apache_beam.runners.worker import bundle_processor from apache_beam.runners.worker import operations +from apache_beam.runners.worker import sdk_worker from apache_beam.runners.worker.bundle_processor import BeamTransformFactory from apache_beam.runners.worker.bundle_processor import BundleProcessor from apache_beam.runners.worker.bundle_processor import DataInputOperation @@ -429,6 +431,116 @@ def test_user_modified_sdks_need_to_be_installed_in_runtime_env(self): "beam:version:sdk_base:apache/beam_python3.5_sdk:2.1.0-custom")) +class RuntimeStateCommitTest(unittest.TestCase): + STATE_TYPES = ( + bundle_processor.SynchronousBagRuntimeState, + bundle_processor.SynchronousSetRuntimeState, + SynchronousOrderedListRuntimeState) + + def setUp(self): + # Keep set compaction deterministic; it is exercised separately below. + patcher = mock.patch.object(random, 'random', return_value=0) + patcher.start() + self.addCleanup(patcher.stop) + + def _create_state(self, state_type, append_errors=('', ), clear_error=''): + futures = [] + + def response_future(error): + future = sdk_worker._Future().set( + beam_fn_api_pb2.StateResponse(error=error)) + future.get = mock.Mock(wraps=future.get) + futures.append(future) + return future + + append_errors = iter(append_errors) + underlying = mock.Mock(spec=sdk_worker.StateHandler) + underlying.append_raw.side_effect = ( + lambda *args: response_future(next(append_errors))) + underlying.clear.side_effect = lambda *args: response_future(clear_error) + underlying.get_raw.return_value = (b'', None) + handler = GlobalCachingStateHandler(StateCache(0), underlying) + if state_type is SynchronousOrderedListRuntimeState: + key = beam_fn_api_pb2.StateKey( + ordered_list_user_state=beam_fn_api_pb2.StateKey.OrderedListUserState( + )) + else: + key = beam_fn_api_pb2.StateKey( + bag_user_state=beam_fn_api_pb2.StateKey.BagUserState()) + return state_type(handler, key, StrUtf8Coder()), futures + + def _add(self, state, value): + if isinstance(state, SynchronousOrderedListRuntimeState): + state.add((timestamp.Timestamp(1), value)) + else: + state.add(value) + + def test_commit_rejects_failed_append(self): + for state_type in self.STATE_TYPES: + with self.subTest(state_type=state_type): + state, _ = self._create_state( + state_type, append_errors=('append failed', )) + self._add(state, 'value') + with self.assertRaisesRegex(RuntimeError, 'append failed'): + state.commit() + + def test_commit_rejects_failed_clear(self): + for state_type in self.STATE_TYPES: + with self.subTest(state_type=state_type): + state, _ = self._create_state(state_type, clear_error='clear failed') + state.clear() + with self.assertRaisesRegex(RuntimeError, 'clear failed'): + state.commit() + + def test_clear_then_append_checks_every_response(self): + for state_type in self.STATE_TYPES: + for clear_error, append_error in ( + ('', ''), ('clear failed', ''), ('clear failed', 'append failed')): + with self.subTest(state_type=state_type, + clear_error=clear_error, + append_error=append_error): + state, futures = self._create_state( + state_type, (append_error, ), clear_error) + state.clear() + self._add(state, 'replacement') + if clear_error: + with self.assertRaises(RuntimeError) as raised: + state.commit() + self.assertEqual( + str(raised.exception), + '\n'.join( + error for error in (clear_error, append_error) if error)) + else: + state.commit() + self.assertEqual(len(futures), 2) + for future in futures: + future.get.assert_called_once() + + def test_commit_rejects_failure_in_earlier_append_chunk(self): + state, futures = self._create_state( + bundle_processor.SynchronousBagRuntimeState, + append_errors=('first append failed', '')) + state.add('first') + state.add('second') + with mock.patch.object(sdk_worker.data_plane, + '_DEFAULT_SIZE_FLUSH_THRESHOLD', + 1): + with self.assertRaisesRegex(RuntimeError, 'first append failed'): + state.commit() + self.assertEqual(len(futures), 2) + for future in futures: + future.get.assert_called_once() + + def test_commit_rejects_failed_set_compaction(self): + state, futures = self._create_state( + bundle_processor.SynchronousSetRuntimeState, clear_error='clear failed') + with mock.patch.object(random, 'random', return_value=1): + state.add('value') + with self.assertRaisesRegex(RuntimeError, 'clear failed'): + state.commit() + self.assertEqual(len(futures), 2) + + class OrderedListStateTest(unittest.TestCase): class NoStateCache(StateCache): def __init__(self): From 7e4aa30f17d28725ccece5ff3b8f912d4c94fc8c Mon Sep 17 00:00:00 2001 From: bvolpato Date: Thu, 3 Sep 2026 18:06:24 +0000 Subject: [PATCH 2/2] [Python] Link state fix in changelog --- CHANGES.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/CHANGES.md b/CHANGES.md index e5d0f727abc5..b55669143bb4 100644 --- a/CHANGES.md +++ b/CHANGES.md @@ -93,7 +93,7 @@ ## Bugfixes * (Java) Fixed the Spark runner firing processing-time timers in reverse timestamp order ([#39824](https://github.com/apache/beam/issues/39824)). -* (Python) State write failures reported by portable runners now fail the bundle instead of silently losing state updates. +* (Python) State write failures reported by portable runners now fail the bundle instead of silently losing state updates ([#39992](https://github.com/apache/beam/pull/39992)). * (Python) Fixed incorrect profiler options handling on portable runners ([#39613](https://github.com/apache/beam/issues/39613)). * (Java) KafkaIO dynamic reads no longer require the obsolete `beam_fn_api` experiment ([#29998](https://github.com/apache/beam/issues/29998)). * (Prism) Self-checkpointing splittable DoFns now resume after their requested delay instead of immediately, so polling SDFs no longer busy-spin ([#39848](https://github.com/apache/beam/issues/39848)).