diff --git a/sdks/python/apache_beam/runners/interactive/interactive_beam_test.py b/sdks/python/apache_beam/runners/interactive/interactive_beam_test.py index 21163fc121c5..af3f6638fc91 100644 --- a/sdks/python/apache_beam/runners/interactive/interactive_beam_test.py +++ b/sdks/python/apache_beam/runners/interactive/interactive_beam_test.py @@ -828,7 +828,7 @@ def test_compute_non_blocking_ipython_widgets( ]) mock_clear_output.assert_called_once() - async_result.result(timeout=60) # Let it finish + async_result.wait_for_completion(timeout=60) def test_compute_dependency_wait_true(self): p = beam.Pipeline(ir.InteractiveRunner()) @@ -853,12 +853,12 @@ def test_compute_dependency_wait_true(self): spy_wait.assert_called_with({pcoll2}, async_res2) # Let pcoll1 finish - async_res1.result(timeout=60) + async_res1.wait_for_completion(timeout=60) self.assertTrue(pcoll1 in self.env.computed_pcollections) self.assertFalse(self.env.is_pcollection_computing(pcoll1)) # pcoll2 should now run and complete - async_res2.result(timeout=60) + async_res2.wait_for_completion(timeout=60) self.assertTrue(pcoll2 in self.env.computed_pcollections) @patch.object(ie.InteractiveEnvironment, 'is_pcollection_computing') @@ -878,7 +878,7 @@ def test_compute_dependency_wait_false(self, mock_is_computing): '_execute_pipeline_fragment', wraps=rm._execute_pipeline_fragment) as spy_execute: async_res2 = ib.compute(pcoll2, blocking=False, wait_for_inputs=False) - async_res2.result(timeout=60) + async_res2.wait_for_completion(timeout=60) # Assert that execute was called for pcoll2 without waiting spy_execute.assert_called_with({pcoll2}, async_res2, ANY, ANY) diff --git a/sdks/python/apache_beam/runners/interactive/recording_manager.py b/sdks/python/apache_beam/runners/interactive/recording_manager.py index c19b60b64fd2..5d242c63f2bf 100644 --- a/sdks/python/apache_beam/runners/interactive/recording_manager.py +++ b/sdks/python/apache_beam/runners/interactive/recording_manager.py @@ -25,6 +25,7 @@ from concurrent.futures import Future from concurrent.futures import ThreadPoolExecutor from typing import Any +from typing import Iterable from typing import Optional from typing import Union @@ -86,6 +87,7 @@ def __init__( bar_style='info', ) if IS_IPYTHON else None) self._cancel_requested = False + self._completed_event = threading.Event() if IS_IPYTHON: self._cancel_button.on_click(self._cancel_clicked) @@ -151,27 +153,40 @@ def exception(self, timeout=None): except TimeoutError: return None - def _on_done(self, future: Future): - self._env.unmark_pcollection_computing(self._pcolls) - self._recording_manager._async_computations.pop(self._display_id, None) - - if future.cancelled(): - self.update_display('Computation Cancelled.', 1.0) - return - - exc = future.exception() + def wait_for_completion(self, timeout=None): + if not self._completed_event.wait(timeout=timeout): + raise TimeoutError( + 'Timeout waiting for asynchronous computation completion.') + if self._future.cancelled(): + raise RuntimeError('Asynchronous computation was cancelled.') + exc = self.exception() if exc: - self.update_display(f'Error: {exc}', 1.0) - _LOGGER.error('Asynchronous computation failed: %s', exc, exc_info=exc) - else: - self.update_display('Computation Finished Successfully.', 1.0) - res = future.result() - if res and res.state == PipelineState.DONE: - self._env.mark_pcollection_computed(self._pcolls) + raise exc + + def _on_done(self, future: Future): + try: + if future.cancelled(): + self.update_display('Computation Cancelled.', 1.0) + return + + exc = future.exception() + if exc: + self.update_display(f'Error: {exc}', 1.0) + _LOGGER.error('Asynchronous computation failed: %s', exc, exc_info=exc) else: - _LOGGER.warning( - 'Async computation finished but state is not DONE: %s', - res.state if res else 'Unknown') + self.update_display('Computation Finished Successfully.', 1.0) + res = future.result() + if res and res.state == PipelineState.DONE: + self._env.mark_pcollection_computed(self._pcolls) + else: + _LOGGER.warning( + 'Async computation finished but state is not DONE: %s', + res.state if res else 'Unknown') + finally: + self._env.unmark_pcollection_computing(self._pcolls) + with self._recording_manager._lock: + self._recording_manager._async_computations.pop(self._display_id, None) + self._completed_event.set() def cancel(self): if self._future.done(): @@ -438,7 +453,9 @@ def __init__( self._executor = ThreadPoolExecutor(max_workers=os.cpu_count()) self._env = ie.current_env() self._async_computations: dict[str, AsyncComputationResult] = {} + self._lock = threading.Lock() self._pipeline_graph = None + self._applied_labels_snapshot = set() def _execute_pipeline_fragment( self, @@ -499,24 +516,10 @@ def _run_async_computation( pipeline_result = self._execute_pipeline_fragment( pcolls_to_compute, async_result, runner, options) - # if pipeline_result.state == PipelineState.DONE: - # self._env.mark_pcollection_computed(pcolls_to_compute) - # _LOGGER.info( - # 'Asynchronous computation finished successfully for' - # f' {len(pcolls_to_compute)} PCollections.' - # ) - # else: - # _LOGGER.error( - # 'Asynchronous computation failed for' - # f' {len(pcolls_to_compute)} PCollections. State:' - # f' {pipeline_result.state}' - # ) return pipeline_result except Exception as e: _LOGGER.exception('Exception during asynchronous computation: %s', e) raise - # finally: - # self._env.unmark_pcollection_computing(pcolls_to_compute) def _watch(self, pcolls: list[beam.pvalue.PCollection]) -> None: """Watch any pcollections not being watched. @@ -575,7 +578,7 @@ def clear(self) -> None: if cache_manager: cache_manager.cleanup() - def cancel(self: None) -> None: + def cancel(self) -> None: """Cancels the current background recording job.""" bcj.attempt_to_cancel_background_caching_job(self.user_pipeline) @@ -694,7 +697,8 @@ def compute_async( future = Future() async_result = AsyncComputationResult( future, pcolls_to_compute, self.user_pipeline, self) - self._async_computations[async_result._display_id] = async_result + with self._lock: + self._async_computations[async_result._display_id] = async_result self._env.mark_pcollection_computing(pcolls_to_compute) def task(): @@ -710,18 +714,22 @@ def task(): return async_result def _get_pipeline_graph(self): - """Lazily initializes and returns the PipelineGraph.""" - if self._pipeline_graph is None: + """Lazily initializes and returns the PipelineGraph, rebuilding it + only if the pipeline transforms have changed. + """ + if (self._pipeline_graph is None or + self._applied_labels_snapshot != self.user_pipeline.applied_labels): try: # Try to create the graph. self._pipeline_graph = PipelineGraph(self.user_pipeline) + self._applied_labels_snapshot = set(self.user_pipeline.applied_labels) except (ImportError, NameError, AttributeError): # If pydot is missing, PipelineGraph() might crash. _LOGGER.warning( - "Could not create PipelineGraph (pydot missing?). " \ - "Async features disabled." - ) + "Could not create PipelineGraph (pydot missing?). " + "Async features disabled.") self._pipeline_graph = None + self._applied_labels_snapshot = set() return self._pipeline_graph def _get_pcoll_id_map(self): @@ -800,11 +808,19 @@ def _wait_for_dependencies( """Waits for any dependencies of the given PCollections that are currently being computed.""" dependencies = self._get_all_dependencies(pcolls) + if async_result is None: + pcolls_to_check = dependencies.union(pcolls) + else: + pcolls_to_check = dependencies computing_deps: dict[beam.pvalue.PCollection, AsyncComputationResult] = {} - for dep in dependencies: - if self._env.is_pcollection_computing(dep): - for comp in self._async_computations.values(): + with self._lock: + async_computations_copy = list(self._async_computations.values()) + + for dep in pcolls_to_check: + is_computing = self._env.is_pcollection_computing(dep) + if is_computing: + for comp in async_computations_copy: if dep in comp._pcolls: computing_deps[dep] = comp break @@ -820,17 +836,16 @@ def _wait_for_dependencies( len(computing_deps), computing_deps.keys()) - futures_to_wait = list( - set(comp._future for comp in computing_deps.values())) + results_to_wait = list(set(comp for comp in computing_deps.values())) try: - for i, future in enumerate(futures_to_wait): + for i, comp in enumerate(results_to_wait): if async_result: async_result.update_display( - f'Waiting for dependency {i + 1}/{len(futures_to_wait)}...', - progress=0.05 + 0.05 * (i / len(futures_to_wait)), + f'Waiting for dependency {i + 1}/{len(results_to_wait)}...', + progress=0.05 + 0.05 * (i / len(results_to_wait)), ) - future.result() + comp.wait_for_completion() if async_result: async_result.update_display('Dependencies finished.', progress=0.1) _LOGGER.info('Dependencies finished successfully.') @@ -841,6 +856,18 @@ def _wait_for_dependencies( _LOGGER.error('Dependency computation failed: %s', e, exc_info=e) return False + def _get_uncomputed_pcolls( + self, + pcolls: Iterable[beam.pvalue.PCollection], + ) -> set[beam.pvalue.PCollection]: + """Filter out already computed PCollections from the environment.""" + current_env = ie.current_env() + computed = { + pcoll + for pcoll in pcolls if pcoll in current_env.computed_pcollections + } + return set(pcolls).difference(computed) + def record( self, pcolls: list[beam.pvalue.PCollection], @@ -878,38 +905,46 @@ def record( self._watch(pcolls) self.record_pipeline() - # Get the subset of computed PCollections. These do not to be recomputed. - computed_pcolls = set( - pcoll for pcoll in pcolls - if pcoll in ie.current_env().computed_pcollections) - - # Start a pipeline fragment to start computing the PCollections. - uncomputed_pcolls = set(pcolls).difference(computed_pcolls) - if uncomputed_pcolls: - if not self._wait_for_dependencies(uncomputed_pcolls): - raise RuntimeError( - 'Cannot record because a dependency failed to compute' - ' asynchronously.') - - self._clear() - - merged_options = pipeline_options.PipelineOptions( - **{ - **self.user_pipeline.options.get_all_options( - drop_default=True, retain_unknown_options=True), - **options.get_all_options( - drop_default=True, retain_unknown_options=True) - }) if options else self.user_pipeline.options - - cache_path = ie.current_env().options.cache_root - is_remote_run = cache_path and ie.current_env( - ).options.cache_root.startswith('gs://') - pf.PipelineFragment( - list(uncomputed_pcolls), merged_options, - runner=runner).run(blocking=is_remote_run) - result = ie.current_env().pipeline_result(self.user_pipeline) - else: - result = None + # Early check and return if everything is already computed + uncomputed_pcolls = self._get_uncomputed_pcolls(pcolls) + if not uncomputed_pcolls: + recording = Recording( + self.user_pipeline, pcolls, None, max_n, max_duration_secs) + self._recordings.add(recording) + return recording + + # Wait for dependencies if there are uncomputed PCollections + if not self._wait_for_dependencies(uncomputed_pcolls): + raise RuntimeError( + 'Cannot record because a dependency failed to compute' + ' asynchronously.') + + # Re-evaluate uncomputed PCollections + uncomputed_pcolls = self._get_uncomputed_pcolls(pcolls) + if not uncomputed_pcolls: + recording = Recording( + self.user_pipeline, pcolls, None, max_n, max_duration_secs) + self._recordings.add(recording) + return recording + + # Flattened execution path (no indentation needed) + self._clear() + + merged_options = pipeline_options.PipelineOptions( + **{ + **self.user_pipeline.options.get_all_options( + drop_default=True, retain_unknown_options=True), + **options.get_all_options( + drop_default=True, retain_unknown_options=True) + }) if options else self.user_pipeline.options + + cache_path = ie.current_env().options.cache_root + is_remote_run = cache_path and ie.current_env( + ).options.cache_root.startswith('gs://') + pf.PipelineFragment( + list(uncomputed_pcolls), merged_options, + runner=runner).run(blocking=is_remote_run) + result = ie.current_env().pipeline_result(self.user_pipeline) recording = Recording( self.user_pipeline, pcolls, result, max_n, max_duration_secs) diff --git a/sdks/python/apache_beam/runners/interactive/recording_manager_test.py b/sdks/python/apache_beam/runners/interactive/recording_manager_test.py index d2038719f67a..55c3bd91cfd6 100644 --- a/sdks/python/apache_beam/runners/interactive/recording_manager_test.py +++ b/sdks/python/apache_beam/runners/interactive/recording_manager_test.py @@ -937,20 +937,22 @@ def test_record_detects_remote_runner( p = beam.Pipeline(InteractiveRunner()) numbers = p | 'numbers' >> beam.Create([0, 1, 2]) - # Set the cache directory for Interactive Beam to be in a GCS bucket. - ib.options.cache_root = 'gs://test-bucket/' + original_cache_root = ib.options.cache_root + try: + # Set the cache directory for Interactive Beam to be in a GCS bucket. + ib.options.cache_root = 'gs://test-bucket/' - # Create the recording objects. By calling `record` a new PipelineFragment - # is started to compute the given PCollections and cache to disk. - rm = RecordingManager(p) + # Create the recording objects. By calling `record` a new PipelineFragment + # is started to compute the given PCollections and cache to disk. + rm = RecordingManager(p) - # Run record() and check if the PipelineFragment.run had blocking set to - # True due to the GCS cache_root value. - rm.record([numbers], max_n=3, max_duration=500) - mock_pipeline_fragment.assert_called_with(blocking=True) + # Run record() and check if the PipelineFragment.run had blocking set to + # True due to the GCS cache_root value. + rm.record([numbers], max_n=3, max_duration=500) + mock_pipeline_fragment.assert_called_with(blocking=True) - # Reset cache_root value. - ib.options.cache_root = None + finally: + ib.options.cache_root = original_cache_root def test_compute_async_blocking(self): p = beam.Pipeline(InteractiveRunner()) @@ -1062,9 +1064,171 @@ def test_wait_for_dependencies(self, mock_async_result_cls): ie.current_env().mark_pcollection_computing({p1}) self.assertTrue(rm._wait_for_dependencies({p2})) - mock_future.result.assert_called_once() + mock_async_res.wait_for_completion.assert_called_once() ie.current_env().unmark_pcollection_computing({p1}) + def test_wait_for_dependencies_async_result(self): + p = beam.Pipeline(InteractiveRunner()) + pcoll = p | beam.Create([1]) + rm = RecordingManager(p) + + # Mock a background computation for pcoll + async_res = MagicMock(spec=AsyncComputationResult) + async_res._pcolls = {pcoll} + rm._async_computations['id1'] = async_res + ie.current_env().mark_pcollection_computing({pcoll}) + + # Case 1: Called from main thread (async_result=None) + # It should wait for pcoll's background job completion event + with patch.object(async_res, 'wait_for_completion') as mock_wait: + with patch.object(rm, '_get_all_dependencies', return_value=set()): + rm._wait_for_dependencies({pcoll}, async_result=None) + mock_wait.assert_called_once() + + # Case 2: Called from background thread (async_result=async_res) + # It should NOT wait (prevents self-deadlock) + with patch.object(async_res, 'wait_for_completion') as mock_wait: + with patch.object(rm, '_get_all_dependencies', return_value=set()): + rm._wait_for_dependencies({pcoll}, async_result=async_res) + mock_wait.assert_not_called() + + def test_record_recalculates_uncomputed_pcolls_after_wait(self): + from apache_beam.runners.interactive import pipeline_fragment as pf + p = beam.Pipeline(InteractiveRunner()) + pcoll = p | beam.Create([1]) + rm = RecordingManager(p) + + # Mock wait_for_dependencies to simulate pcoll completing during the wait + def mock_wait(pcolls, async_result=None): + # Simulate pcoll completing by marking it computed + ie.current_env().mark_pcollection_computed({pcoll}) + return True + + with patch.object(rm, '_wait_for_dependencies', side_effect=mock_wait), \ + patch.object(rm, '_clear') as mock_clear, \ + patch.object(pf, 'PipelineFragment') as mock_fragment: + + rm.record([pcoll], max_n=10, max_duration='inf') + + # Since pcoll was marked computed during the wait, + # uncomputed_pcolls becomes empty. + # So _clear() and PipelineFragment should NOT be called. + mock_clear.assert_not_called() + mock_fragment.assert_not_called() + + def test_get_pipeline_graph_not_cached(self): + p = beam.Pipeline(InteractiveRunner()) + pcoll1 = p | 'Create1' >> beam.Create([1]) + rm = RecordingManager(p) + + # First call: graph is created + graph1 = rm._get_pipeline_graph() + self.assertIsNotNone(graph1) + self.assertEqual(graph1._pipeline_instrument.user_pipeline, p) + + # Add a new transform to the pipeline + pcoll2 = pcoll1 | 'Map1' >> beam.Map(lambda x: x + 1) + + # Second call: graph should be updated and contain the new Map1 transform + graph2 = rm._get_pipeline_graph() + self.assertIsNotNone(graph2) + self.assertIsNot(graph1, graph2) + + # Verify that the new graph contains the newly added transform + transform_names = [ + t.unique_name + for t in graph2._pipeline_proto.components.transforms.values() + ] + self.assertIn('Map1', transform_names) + + def test_wait_for_completion_raises_exception_on_failure(self): + future = Future() + async_res = AsyncComputationResult( + future, set(), beam.Pipeline(InteractiveRunner()), MagicMock()) + # Set the future exception + exc = ValueError("computation error") + future.set_exception(exc) + async_res._completed_event.set() # Simulate completion + with self.assertRaises(ValueError) as context: + async_res.wait_for_completion(timeout=60) + self.assertEqual(context.exception, exc) + + def test_wait_for_completion_raises_error_on_cancellation(self): + future = Future() + async_res = AsyncComputationResult( + future, set(), beam.Pipeline(InteractiveRunner()), MagicMock()) + future.cancel() + async_res._completed_event.set() + with self.assertRaises(RuntimeError) as context: + async_res.wait_for_completion(timeout=60) + self.assertIn('computation was cancelled', str(context.exception)) + + def test_wait_for_completion_raises_timeout_error(self): + future = Future() + async_res = AsyncComputationResult( + future, set(), beam.Pipeline(InteractiveRunner()), MagicMock()) + # completed_event is NOT set, so wait will timeout + with self.assertRaises(TimeoutError): + async_res.wait_for_completion(timeout=0.01) + + def test_get_uncomputed_pcolls(self): + p = beam.Pipeline(InteractiveRunner()) + p1 = p | 'C1' >> beam.Create([1]) + p2 = p | 'C2' >> beam.Create([2]) + ib.watch(locals()) + ie.current_env().track_user_pipelines() + + rm = RecordingManager(p) + + # Initially both are uncomputed + self.assertEqual(rm._get_uncomputed_pcolls([p1, p2]), {p1, p2}) + + # Mark p1 as computed + ie.current_env().mark_pcollection_computed([p1]) + self.assertEqual(rm._get_uncomputed_pcolls([p1, p2]), {p2}) + + # Cleanup + ie.current_env().evict_computed_pcollections() + + def test_get_pipeline_graph_caching(self): + p = beam.Pipeline(InteractiveRunner()) + p1 = p | 'C1' >> beam.Create([1]) + ib.watch(locals()) + ie.current_env().track_user_pipelines() + + rm = RecordingManager(p) + graph1 = rm._get_pipeline_graph() + graph2 = rm._get_pipeline_graph() + + # Graph instance is cached and reused when pipeline transforms have not changed + self.assertIs(graph1, graph2) + + # Applying a new transform updates user_pipeline.applied_labels + _ = p1 | 'M1' >> beam.Map(lambda x: x) + graph3 = rm._get_pipeline_graph() + + # Graph is rebuilt because pipeline applied_labels changed + self.assertIsNot(graph1, graph3) + + def test_record_early_return_when_all_computed(self): + p = beam.Pipeline(InteractiveRunner()) + p1 = p | 'C1' >> beam.Create([1]) + ib.watch(locals()) + ie.current_env().track_user_pipelines() + + rm = RecordingManager(p) + ie.current_env().mark_pcollection_computed([p1]) + + with patch.object(rm, '_execute_pipeline_fragment') as mock_exec: + recording = rm.record([p1], max_n=10, max_duration=100) + # Fragment execution must NOT be called + mock_exec.assert_not_called() + self.assertIsNone(recording._result) + self.assertEqual( + recording.wait_until_finish(), beam.runners.runner.PipelineState.DONE) + + ie.current_env().evict_computed_pcollections() + if __name__ == '__main__': unittest.main()