diff --git a/sqldol/sql_base.py b/sqldol/sql_base.py index 5c79979..bd2723f 100644 --- a/sqldol/sql_base.py +++ b/sqldol/sql_base.py @@ -98,9 +98,11 @@ 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``, 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) - ).first()[0] + ).fetchone()[0] @lazyprop def _row_count(self): @@ -151,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: @@ -437,30 +439,55 @@ 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. + ``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:") + >>> _ = 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 batch_size < 1: + raise ValueError(f"batch_size must be a positive integer, got {batch_size}") + if offset < 0 or limit < 0: + raise ValueError( + f"offset and limit must be non-negative, got offset={offset}, limit={limit}" ) - if r.rowcount: - for i, x in enumerate(r.fetchall(), i): - if i < limit: - yield x - else: - break - else: - break + + 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 new file mode 100644 index 0000000..f52c03c --- /dev/null +++ b/sqldol/tests/test_legacy_raw_sql_rows.py @@ -0,0 +1,106 @@ +"""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 + + +@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): + 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]] + assert list(rows[2:2]) == [] + assert list(rows[:0]) == [] + + +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()