diff --git a/sdks/python/apache_beam/transforms/util.py b/sdks/python/apache_beam/transforms/util.py index e5fcf369342c..d99ebb456b56 100644 --- a/sdks/python/apache_beam/transforms/util.py +++ b/sdks/python/apache_beam/transforms/util.py @@ -1204,7 +1204,7 @@ def finish_bundle(self): # Check if adding this element would exceed limits would_exceed_count = len(batch) >= self._max_batch_size would_exceed_weight = ( - batch_weight + element_size >= self._max_batch_weight and batch) + batch_weight + element_size > self._max_batch_weight and batch) if would_exceed_count or would_exceed_weight: # Emit current batch @@ -1301,7 +1301,7 @@ def _flush_window(self, win): would_exceed_count = len(batch) >= self._max_batch_size would_exceed_weight = ( - batch_weight + element_size >= self._max_batch_weight and batch) + batch_weight + element_size > self._max_batch_weight and batch) if would_exceed_count or would_exceed_weight: yield windowed_value.WindowedValue(batch, win.max_timestamp(), (win, )) diff --git a/sdks/python/apache_beam/transforms/util_test.py b/sdks/python/apache_beam/transforms/util_test.py index 5a935731ca71..e8e7d762b819 100644 --- a/sdks/python/apache_beam/transforms/util_test.py +++ b/sdks/python/apache_beam/transforms/util_test.py @@ -1416,6 +1416,22 @@ def test_global_dofn_weight_splitting(self): for batch in batches: self.assertEqual(len(batch), 2) + def test_global_dofn_batch_can_reach_max_batch_weight(self): + """Test that a batch can weigh exactly max_batch_weight.""" + from apache_beam.transforms.util import _SortAndBatchElementsDoFn + + # Each element has size 5, max_batch_weight=10 -> 2 per batch + dofn = _SortAndBatchElementsDoFn( + min_batch_size=1, + max_batch_size=100, + max_batch_weight=10, + element_size_fn=len) + dofn.start_bundle() + for elem in ['aaaaa', 'bbbbb', 'ccccc', 'ddddd']: + dofn.process(elem) + batches = [wv.value for wv in dofn.finish_bundle()] + self.assertEqual([len(batch) for batch in batches], [2, 2]) + def test_windowed_dofn_flush_and_finish(self): """Test _WindowAwareSortAndBatchElementsDoFn directly.""" from apache_beam.transforms.util import _WindowAwareSortAndBatchElementsDoFn @@ -1493,6 +1509,21 @@ def test_windowed_dofn_weight_splitting(self): self.assertEqual(len(wv.value), 2) self.assertEqual(wv.windows[0], win) + def test_windowed_dofn_batch_can_reach_max_batch_weight(self): + """Test that a windowed batch can weigh exactly max_batch_weight.""" + from apache_beam.transforms.util import _WindowAwareSortAndBatchElementsDoFn + + dofn = _WindowAwareSortAndBatchElementsDoFn( + min_batch_size=1, + max_batch_size=100, + max_batch_weight=10, + element_size_fn=len) + dofn.start_bundle() + win = IntervalWindow(0, 10) + dofn._buffers[win].extend(['aaaaa', 'bbbbb', 'ccccc', 'ddddd']) + batches = list(dofn._flush_window(win)) + self.assertEqual([len(wv.value) for wv in batches], [2, 2]) + class IdentityWindowTest(unittest.TestCase): def test_window_preserved(self):