From 599f68ba1dfd400efd243047d9d094ef3b2ef895 Mon Sep 17 00:00:00 2001 From: Thor Whalen <1906276+thorwhalen@users.noreply.github.com> Date: Tue, 22 Sep 2026 15:25:22 +0000 Subject: [PATCH 1/3] fix(sql_base): iter_rows detects the end of a table; count_rows on DB-API - iter_rows decided whether a page had rows from the cursor's rowcount, which is -1 (truthy) for a SELECT on sqlite3, so without a limit it requested empty pages until limit (default 1e12) ran out. It now stops on the first page shorter than requested. It also no longer yields more than limit rows when limit spans several pages (the running counter restarted at the last index), and a page never asks for more rows than remain under limit. A batch_size < 1 raises a ValueError naming it. - SqlTableRowsCollection.count_rows used .first(), a SQLAlchemy result method a DB-API cursor lacks; it now uses fetchone(), which both have. Tests: sqldol/tests/test_legacy_raw_sql_rows.py (sqlite3, with a query-counting connection so the old code fails instead of hanging). The workaround note in test_sql_injection.py is dropped. Co-Authored-By: Claude Opus 5 --- sqldol/sql_base.py | 52 ++++++++----- sqldol/tests/test_legacy_raw_sql_rows.py | 94 ++++++++++++++++++++++++ sqldol/tests/test_sql_injection.py | 4 +- 3 files changed, 129 insertions(+), 21 deletions(-) create mode 100644 sqldol/tests/test_legacy_raw_sql_rows.py diff --git a/sqldol/sql_base.py b/sqldol/sql_base.py index 5c79979..6665120 100644 --- a/sqldol/sql_base.py +++ b/sqldol/sql_base.py @@ -98,9 +98,12 @@ def __init__(self, connection, table_name, batch_size=2000, limit=int(1e16)): # QUESTION: Should helpers (_describe, _columns, etc.) should be methods/properties/lazyprops, and hidden or not? def count_rows(self): + # Note: ``fetchone`` is DB-API, so this works on a plain DB-API connection + # (whose ``execute`` returns a cursor, which has no ``first``) as well as on + # a SQLAlchemy result. return self.connection.execute( self._tmpl_count_rows_tmpl.format(table_name=self.table_name) - ).first()[0] + ).fetchone()[0] @lazyprop def _row_count(self): @@ -437,30 +440,43 @@ def __len__(self): def iter_rows(connection, table_name, batch_size=1000, offset=0, limit=int(1e12)): - """Iterate the over the rows of a table. - The limit argument is mostly there to avoid an infinite loop, but can also be used to get ranges. + """Iterate over the rows of a table, fetching ``batch_size`` rows per query. + + Yields at most ``limit`` rows, starting at row ``offset`` (like SQL's + ``LIMIT``/``OFFSET``), and stops at the end of the table: a page shorter than + the one requested means there is nothing after it. ``table_name`` must pass :func:`validate_sql_identifier`, and ``batch_size``, ``offset`` and ``limit`` must be integers, because they are written into the SQL text. + + >>> import sqlite3 + >>> con = sqlite3.connect(":memory:") + >>> _ = con.execute("CREATE TABLE t (x INTEGER)") + >>> _ = con.executemany("INSERT INTO t VALUES (?)", [(i,) for i in range(5)]) + >>> list(iter_rows(con, "t", batch_size=2)) + [(0,), (1,), (2,), (3,), (4,)] + >>> list(iter_rows(con, "t", batch_size=2, offset=1, limit=3)) + [(1,), (2,), (3,)] """ table_name = validate_sql_identifier(table_name) batch_size, offset, limit = index(batch_size), index(offset), index(limit) - stop_offset = limit + offset # to the range act like like sql limit - i = 0 - for offset in range(offset, stop_offset, batch_size): - if i >= limit: - break - r = connection.execute( - f"SELECT * FROM {table_name} LIMIT {batch_size} OFFSET {offset}" - ) - if r.rowcount: - for i, x in enumerate(r.fetchall(), i): - if i < limit: - yield x - else: - break - else: + if batch_size < 1: + raise ValueError(f"batch_size must be a positive integer, got {batch_size}") + # Note: The end of the table is detected from the rows actually fetched, never + # from ``rowcount``: for a SELECT that is -1 on DB-API drivers that don't + # pre-buffer (sqlite3), so testing it made this loop request empty pages until + # ``limit`` ran out. + remaining = limit + while remaining > 0: + page_size = min(batch_size, remaining) + rows = connection.execute( + f"SELECT * FROM {table_name} LIMIT {page_size} OFFSET {offset}" + ).fetchall() + yield from rows + if len(rows) < page_size: break + remaining -= page_size + offset += page_size ####### Stores ######################################################################################################## diff --git a/sqldol/tests/test_legacy_raw_sql_rows.py b/sqldol/tests/test_legacy_raw_sql_rows.py new file mode 100644 index 0000000..1846ea4 --- /dev/null +++ b/sqldol/tests/test_legacy_raw_sql_rows.py @@ -0,0 +1,94 @@ +"""The legacy raw-SQL row readers (``sqldol.sql_base``) on a plain DB-API connection. + +``iter_rows`` used to decide whether a page had rows from the cursor's +``rowcount``, which is -1 (truthy) for a SELECT on sqlite3, so without a small +``limit`` it kept requesting empty pages. It also over-yielded when ``limit`` +spanned several pages. ``SqlTableRowsCollection.count_rows`` called ``.first()``, +a SQLAlchemy result method that a DB-API cursor does not have. +""" + +import math +import sqlite3 + +import pytest + +from sqldol.sql_base import SqlTableRowsCollection, SqlTableRowsSequence, iter_rows + +TABLE = "t" +N_ROWS = 7 + + +class CountingConnection: + """A sqlite3 connection that counts queries and refuses to run forever.""" + + def __init__(self, connection, *, max_queries=100): + self._connection = connection + self.max_queries = max_queries + self.queries = 0 + + def execute(self, sql, *args): + self.queries += 1 + if self.queries > self.max_queries: + raise AssertionError(f"{self.queries} queries: the end was never detected") + return self._connection.execute(sql, *args) + + +@pytest.fixture +def connection(): + con = sqlite3.connect(":memory:") + con.execute(f"CREATE TABLE {TABLE} (x INTEGER)") + con.executemany(f"INSERT INTO {TABLE} VALUES (?)", [(i,) for i in range(N_ROWS)]) + con.commit() + yield CountingConnection(con) + con.close() + + +ALL_ROWS = [(i,) for i in range(N_ROWS)] + + +@pytest.mark.parametrize("batch_size", [1, 2, 3, N_ROWS, N_ROWS + 5]) +def test_iter_rows_stops_at_the_end_of_the_table(connection, batch_size): + assert list(iter_rows(connection, TABLE, batch_size=batch_size)) == ALL_ROWS + # One query per full page, plus the short (possibly empty) one that ends it. + assert connection.queries == N_ROWS // batch_size + 1 + + +def test_iter_rows_on_an_empty_table_makes_one_query(connection): + connection.execute(f"DELETE FROM {TABLE}") + assert list(iter_rows(connection, TABLE, batch_size=3)) == [] + assert connection.queries == 2 # the DELETE, then one SELECT + + +@pytest.mark.parametrize("batch_size", [1, 2, 3, 10]) +@pytest.mark.parametrize( + "offset, limit", [(0, 0), (0, 1), (0, 5), (2, 3), (2, 100), (6, 4), (7, 2), (9, 1)] +) +def test_iter_rows_offset_and_limit_act_like_a_slice( + connection, batch_size, offset, limit +): + rows = list( + iter_rows(connection, TABLE, batch_size=batch_size, offset=offset, limit=limit) + ) + assert rows == ALL_ROWS[offset : offset + limit] + assert connection.queries <= math.ceil(limit / batch_size) + 1 + + +def test_iter_rows_rejects_a_non_positive_batch_size(connection): + with pytest.raises(ValueError, match="batch_size"): + list(iter_rows(connection, TABLE, batch_size=0)) + + +def test_rows_collection_counts_iterates_and_slices(connection): + rows = SqlTableRowsCollection(connection, TABLE, batch_size=3) + assert rows.count_rows() == N_ROWS + assert len(rows) == N_ROWS + assert list(rows) == ALL_ROWS + assert list(rows[2:5]) == ALL_ROWS[2:5] + assert list(rows[4:]) == ALL_ROWS[4:] + assert list(rows[3]) == [ALL_ROWS[3]] + + +def test_rows_sequence_indexes_like_a_list(connection): + rows = SqlTableRowsSequence(connection, TABLE, batch_size=2) + assert rows[1:6] == ALL_ROWS[1:6] + assert len(rows) == N_ROWS diff --git a/sqldol/tests/test_sql_injection.py b/sqldol/tests/test_sql_injection.py index 4061dcc..8a91ca1 100644 --- a/sqldol/tests/test_sql_injection.py +++ b/sqldol/tests/test_sql_injection.py @@ -250,7 +250,5 @@ def test_iter_rows_rejects_a_table_name_carrying_sql(dbapi_connection): def test_raw_sql_paths_still_read_a_plain_table(dbapi_connection): SqlTableRowsCollection(dbapi_connection, TABLE_NAME) # accepted - # Note: a single bounded batch, because on a DB-API cursor whose SELECT rowcount - # is -1 (sqlite3), iter_rows does not detect the end of the table by itself. - rows = iter_rows(dbapi_connection, TABLE_NAME, batch_size=10, limit=10) + rows = iter_rows(dbapi_connection, TABLE_NAME, batch_size=10) assert sorted(rows) == _original_rows() From 570bbe537d846460d7915779a8087cc87bd580c2 Mon Sep 17 00:00:00 2001 From: Thor Whalen <1906276+thorwhalen@users.noreply.github.com> Date: Tue, 22 Sep 2026 15:28:00 +0000 Subject: [PATCH 2/3] review: reject negative offset/limit eagerly; rows[:0] is empty; fix count_rows note - iter_rows validates batch_size/offset/limit when called (not on first next()); a negative offset duplicated rows on sqlite and a negative limit silently gave [], both now ValueError. - SqlTableRowsCollection[a:0] returned every row (`if stop:`). - count_rows comment no longer claims SQLAlchemy 2 connections work: the raw-SQL strings here only execute on DB-API connections. Co-Authored-By: Claude Opus 5 --- sqldol/sql_base.py | 50 ++++++++++++++---------- sqldol/tests/test_legacy_raw_sql_rows.py | 18 +++++++-- 2 files changed, 45 insertions(+), 23 deletions(-) diff --git a/sqldol/sql_base.py b/sqldol/sql_base.py index 6665120..c78a84c 100644 --- a/sqldol/sql_base.py +++ b/sqldol/sql_base.py @@ -98,9 +98,8 @@ def __init__(self, connection, table_name, batch_size=2000, limit=int(1e16)): # QUESTION: Should helpers (_describe, _columns, etc.) should be methods/properties/lazyprops, and hidden or not? def count_rows(self): - # Note: ``fetchone`` is DB-API, so this works on a plain DB-API connection - # (whose ``execute`` returns a cursor, which has no ``first``) as well as on - # a SQLAlchemy result. + # Note: ``fetchone``, not ``first``: on a DB-API connection (the kind these + # raw-SQL paths run on) ``execute`` returns a cursor, which has no ``first``. return self.connection.execute( self._tmpl_count_rows_tmpl.format(table_name=self.table_name) ).fetchone()[0] @@ -154,7 +153,7 @@ def __getitem__(self, idx): # TODO: Make the ss[:-4] case work too assert step is None, "__getitem__ doesn't handle stepped slices" assert start >= 0, "slice start can't be negative" - if stop: + if stop is not None: assert stop >= start, "slice stop must be at least the slice start" return self.iter_rows(offset=start, limit=stop - start) elif start: @@ -447,7 +446,9 @@ def iter_rows(connection, table_name, batch_size=1000, offset=0, limit=int(1e12) the one requested means there is nothing after it. ``table_name`` must pass :func:`validate_sql_identifier`, and ``batch_size``, - ``offset`` and ``limit`` must be integers, because they are written into the SQL text. + ``offset`` and ``limit`` must be integers, because they are written into the SQL text + (``batch_size`` positive, ``offset`` and ``limit`` non-negative). Invalid arguments + raise ``ValueError`` at call time, before any row is requested. >>> import sqlite3 >>> con = sqlite3.connect(":memory:") @@ -462,21 +463,30 @@ def iter_rows(connection, table_name, batch_size=1000, offset=0, limit=int(1e12) batch_size, offset, limit = index(batch_size), index(offset), index(limit) if batch_size < 1: raise ValueError(f"batch_size must be a positive integer, got {batch_size}") - # Note: The end of the table is detected from the rows actually fetched, never - # from ``rowcount``: for a SELECT that is -1 on DB-API drivers that don't - # pre-buffer (sqlite3), so testing it made this loop request empty pages until - # ``limit`` ran out. - remaining = limit - while remaining > 0: - page_size = min(batch_size, remaining) - rows = connection.execute( - f"SELECT * FROM {table_name} LIMIT {page_size} OFFSET {offset}" - ).fetchall() - yield from rows - if len(rows) < page_size: - break - remaining -= page_size - offset += page_size + if offset < 0 or limit < 0: + raise ValueError( + f"offset and limit must be non-negative, got offset={offset}, limit={limit}" + ) + def pages(): + # Note: The end of the table is detected from the rows actually fetched, + # never from ``rowcount``: for a SELECT that is -1 on DB-API drivers that + # don't pre-buffer (sqlite3), so testing it made this loop request empty + # pages until ``limit`` ran out. + start, remaining = offset, limit + while remaining > 0: + page_size = min(batch_size, remaining) + rows = connection.execute( + f"SELECT * FROM {table_name} LIMIT {page_size} OFFSET {start}" + ).fetchall() + yield from rows + if len(rows) < page_size: + break + remaining -= page_size + start += page_size + + # Note: Not itself a generator, so bad arguments raise when ``iter_rows`` is + # called; the rows are still fetched lazily, page by page. + return pages() ####### Stores ######################################################################################################## diff --git a/sqldol/tests/test_legacy_raw_sql_rows.py b/sqldol/tests/test_legacy_raw_sql_rows.py index 1846ea4..f52c03c 100644 --- a/sqldol/tests/test_legacy_raw_sql_rows.py +++ b/sqldol/tests/test_legacy_raw_sql_rows.py @@ -73,9 +73,19 @@ def test_iter_rows_offset_and_limit_act_like_a_slice( assert connection.queries <= math.ceil(limit / batch_size) + 1 -def test_iter_rows_rejects_a_non_positive_batch_size(connection): - with pytest.raises(ValueError, match="batch_size"): - list(iter_rows(connection, TABLE, batch_size=0)) +@pytest.mark.parametrize( + "kwargs, match", + [ + (dict(batch_size=0), "batch_size"), + (dict(batch_size=-1), "batch_size"), + (dict(offset=-1), "non-negative"), + (dict(limit=-1), "non-negative"), + ], +) +def test_iter_rows_rejects_bad_arguments_when_called(connection, kwargs, match): + with pytest.raises(ValueError, match=match): + iter_rows(connection, TABLE, **kwargs) # not even iterated + assert connection.queries == 0 def test_rows_collection_counts_iterates_and_slices(connection): @@ -86,6 +96,8 @@ def test_rows_collection_counts_iterates_and_slices(connection): assert list(rows[2:5]) == ALL_ROWS[2:5] assert list(rows[4:]) == ALL_ROWS[4:] assert list(rows[3]) == [ALL_ROWS[3]] + assert list(rows[2:2]) == [] + assert list(rows[:0]) == [] def test_rows_sequence_indexes_like_a_list(connection): From a9e0c5b4d8a4a1a9c268ee526226b08d727c1fb1 Mon Sep 17 00:00:00 2001 From: Thor Whalen <1906276+thorwhalen@users.noreply.github.com> Date: Tue, 22 Sep 2026 15:28:05 +0000 Subject: [PATCH 3/3] style: ruff format Co-Authored-By: Claude Opus 5 --- sqldol/sql_base.py | 1 + 1 file changed, 1 insertion(+) diff --git a/sqldol/sql_base.py b/sqldol/sql_base.py index c78a84c..bd2723f 100644 --- a/sqldol/sql_base.py +++ b/sqldol/sql_base.py @@ -467,6 +467,7 @@ def iter_rows(connection, table_name, batch_size=1000, offset=0, limit=int(1e12) raise ValueError( f"offset and limit must be non-negative, got offset={offset}, limit={limit}" ) + def pages(): # Note: The end of the table is detected from the rows actually fetched, # never from ``rowcount``: for a SELECT that is -1 on DB-API drivers that