This is an automated email from the ASF dual-hosted git repository.

shunping pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/beam.git


The following commit(s) were added to refs/heads/master by this push:
     new c5d866c50fe Let SortAndBatchElements batches reach max_batch_weight 
(#40462)
c5d866c50fe is described below

commit c5d866c50fe58e41920090a27fa30575e9281cd1
Author: Anish Mehta <[email protected]>
AuthorDate: Thu Oct 8 19:44:16 2026 +0530

    Let SortAndBatchElements batches reach max_batch_weight (#40462)
    
    A batch was cut when adding the next element would make it equal to 
max_batch_weight, so it never reached the limit. Use > like BatchElements.
    
    fixes #40461
---
 sdks/python/apache_beam/transforms/util.py      |  4 ++--
 sdks/python/apache_beam/transforms/util_test.py | 31 +++++++++++++++++++++++++
 2 files changed, 33 insertions(+), 2 deletions(-)

diff --git a/sdks/python/apache_beam/transforms/util.py 
b/sdks/python/apache_beam/transforms/util.py
index e5fcf369342..d99ebb456b5 100644
--- a/sdks/python/apache_beam/transforms/util.py
+++ b/sdks/python/apache_beam/transforms/util.py
@@ -1204,7 +1204,7 @@ class _SortAndBatchElementsDoFn(DoFn):
       # 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 @@ class _WindowAwareSortAndBatchElementsDoFn(DoFn):
 
       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 5a935731ca7..e8e7d762b81 100644
--- a/sdks/python/apache_beam/transforms/util_test.py
+++ b/sdks/python/apache_beam/transforms/util_test.py
@@ -1416,6 +1416,22 @@ class 
SortAndBatchElementsDoFnDirectTest(unittest.TestCase):
     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 @@ class 
SortAndBatchElementsDoFnDirectTest(unittest.TestCase):
       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):

Reply via email to