From 3e9651cca9c7a605320f4b2bd5638dce78b15e77 Mon Sep 17 00:00:00 2001 From: quettabit <27509167+quettabit@users.noreply.github.com> Date: Thu, 10 Sep 2026 15:21:01 -0700 Subject: [PATCH] initial commit --- src/s2_sdk/_batching.py | 26 ++++++++++++++++---------- src/s2_sdk/_producer.py | 6 ++++++ tests/test_batching.py | 21 ++++++++++++++++++--- 3 files changed, 40 insertions(+), 13 deletions(-) diff --git a/src/s2_sdk/_batching.py b/src/s2_sdk/_batching.py index d89520a..6e1d83f 100644 --- a/src/s2_sdk/_batching.py +++ b/src/s2_sdk/_batching.py @@ -19,6 +19,9 @@ def add(self, record: Record) -> None: self._records.append(record) self._bytes += metered_bytes((record,)) + def would_exceed_max_bytes(self, record: Record) -> bool: + return self._bytes + metered_bytes((record,)) > self._batching.max_bytes + def take(self) -> list[Record]: records = list(self._records) self._records.clear() @@ -60,22 +63,17 @@ async def append_record_batches( next_record_task = None try: - while True: - if next_record_task is not None: - record = await next_record_task - next_record_task = None - else: - record = await anext(record_iter, None) - if record is None: - break - + record = await anext(record_iter, None) + while record is not None: acc.add(record) + record = None deadline = ( asyncio.get_running_loop().time() + linger_secs if linger_secs > 0 else None ) + while not acc.is_full(): if deadline is not None: remaining = deadline - asyncio.get_running_loop().time() @@ -89,11 +87,19 @@ async def append_record_batches( next_record_task = None else: record = await anext(record_iter, None) - if record is None: + if record is None or acc.would_exceed_max_bytes(record): break acc.add(record) + record = None yield acc.take() + + if record is None: + if next_record_task is not None: + record = await next_record_task + next_record_task = None + else: + record = await anext(record_iter, None) finally: if next_record_task is not None: if not next_record_task.done(): diff --git a/src/s2_sdk/_producer.py b/src/s2_sdk/_producer.py index 9bb9732..7a4d747 100644 --- a/src/s2_sdk/_producer.py +++ b/src/s2_sdk/_producer.py @@ -110,6 +110,12 @@ async def submit(self, record: Record) -> RecordSubmitTicket: if self._error is not None: raise self._error + if ( + not self._accumulator.is_empty() + and self._accumulator.would_exceed_max_bytes(record) + ): + await self._submit_batch_now() + loop = asyncio.get_running_loop() ack_fut: asyncio.Future[IndexedAppendAck] = loop.create_future() self._indexed_ack_futs.append(ack_fut) diff --git a/tests/test_batching.py b/tests/test_batching.py index 0a09b14..c52a7a8 100644 --- a/tests/test_batching.py +++ b/tests/test_batching.py @@ -37,8 +37,8 @@ async def test_count_limit(): @pytest.mark.asyncio -async def test_bytes_limit(): - # Each record: 8 bytes overhead + body. Body of 10 bytes → 18 metered bytes. +async def test_batch_flushes_at_max_bytes(): + # Each record is 10(body) + 8(overhead) = 18 metered bytes. records = [Record(body=b"x" * 10) for _ in range(3)] batches = [] async for batch in append_record_batches( @@ -46,12 +46,27 @@ async def test_bytes_limit(): batching=Batching(max_bytes=36, linger=timedelta(0)), ): batches.append(batch) - # 36 bytes limit: first 2 records fit (36 bytes), third goes in next batch assert len(batches) == 2 assert len(batches[0]) == 2 assert len(batches[1]) == 1 +@pytest.mark.asyncio +async def test_batch_flushes_before_exceeding_max_bytes(): + # Each record is 10(body) + 8(overhead) = 18 metered bytes. + records = [Record(body=b"x" * 10) for _ in range(3)] + batches = [] + async for batch in append_record_batches( + _async_iter(records), + batching=Batching(max_bytes=30, linger=timedelta(0)), + ): + batches.append(batch) + assert len(batches) == 3 + assert len(batches[0]) == 1 + assert len(batches[1]) == 1 + assert len(batches[2]) == 1 + + @pytest.mark.asyncio async def test_oversized_record_passes(): records = [Record(body=b"x" * 100)]