|
18 | 18 | """Unit tests for the Beam State and Timer API interfaces.""" |
19 | 19 | # pytype: skip-file |
20 | 20 |
|
| 21 | +import queue |
| 22 | +import threading |
| 23 | +import time |
21 | 24 | import unittest |
22 | 25 | from typing import Any |
23 | 26 |
|
|
34 | 37 | from apache_beam.portability.api import beam_runner_api_pb2 |
35 | 38 | from apache_beam.runners import pipeline_context |
36 | 39 | 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 |
37 | 42 | from apache_beam.testing.test_pipeline import TestPipeline |
38 | 43 | from apache_beam.testing.test_stream import TestStream |
39 | 44 | from apache_beam.testing.util import assert_that |
@@ -719,6 +724,70 @@ def process(self, element, set_state=beam.DoFn.StateParam(SET_STATE)): |
719 | 724 | actual_values = (values | beam.ParDo(SetStatefulDoFn())) |
720 | 725 | assert_that(actual_values, equal_to([1, 3, 6, 10, 10])) |
721 | 726 |
|
| 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 | + |
722 | 791 | def test_stateful_set_state_clean_portably(self): |
723 | 792 | class SetStateClearingStatefulDoFn(beam.DoFn): |
724 | 793 |
|
|
0 commit comments