Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 16 additions & 10 deletions src/s2_sdk/_batching.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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()
Expand All @@ -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():
Expand Down
6 changes: 6 additions & 0 deletions src/s2_sdk/_producer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Comment thread
greptile-apps[bot] marked this conversation as resolved.

loop = asyncio.get_running_loop()
ack_fut: asyncio.Future[IndexedAppendAck] = loop.create_future()
self._indexed_ack_futs.append(ack_fut)
Expand Down
21 changes: 18 additions & 3 deletions tests/test_batching.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,21 +37,36 @@ 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(
_async_iter(records),
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)]
Expand Down
Loading