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
67 changes: 47 additions & 20 deletions sqldol/sql_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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 ########################################################################################################
Expand Down
106 changes: 106 additions & 0 deletions sqldol/tests/test_legacy_raw_sql_rows.py
Original file line number Diff line number Diff line change
@@ -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
4 changes: 1 addition & 3 deletions sqldol/tests/test_sql_injection.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Loading