diff --git a/sqldol/base.py b/sqldol/base.py index a372e19..8f93e7e 100644 --- a/sqldol/base.py +++ b/sqldol/base.py @@ -99,8 +99,14 @@ def __iter__(self): yield from result def __len__(self): - with rows_iter(self.table, self.filt, engine=self.engine) as result: - return result.rowcount + # Note: Counted by the database, not read off a SELECT's ``rowcount``, which + # is -1 on drivers that don't pre-buffer results (SQLite, for one) -- so + # ``len()`` raised ``ValueError`` there instead of counting. + query = select(func.count()).select_from(self.table) + if self.filt is not None: + query = query.where(self.filt) + with self.engine.connect() as connection: + return connection.execute(query).scalar_one() @property def table_name(self): diff --git a/sqldol/tests/test_table_rows_len.py b/sqldol/tests/test_table_rows_len.py new file mode 100644 index 0000000..87a1d3a --- /dev/null +++ b/sqldol/tests/test_table_rows_len.py @@ -0,0 +1,36 @@ +"""Tests for ``TableRows.__len__``: it must count rows, not report a SELECT's rowcount. + +Uses in-memory SQLite only -- no external service, no credentials. +""" + +import pytest +from sqlalchemy import create_engine, text + +from sqldol import TableRows + +TABLE_NAME = "t" + + +@pytest.fixture +def engine(): + """An in-memory SQLite engine holding ``t(k, v)`` with three rows.""" + engine = create_engine("sqlite:///:memory:") + with engine.connect() as connection: + connection.execute(text(f"CREATE TABLE {TABLE_NAME} (k TEXT, v INTEGER)")) + connection.execute( + text(f"INSERT INTO {TABLE_NAME} VALUES ('a', 1), ('b', 2), ('c', 3)") + ) + connection.commit() + return engine + + +def test_len_counts_the_rows(engine): + """``rowcount`` of a SELECT is -1 on SQLite, so ``len`` used to raise ValueError.""" + rows = TableRows(TABLE_NAME, engine=engine) + assert len(rows) == 3 == len(list(rows)) + + +def test_len_honours_the_filter(engine): + rows = TableRows(TABLE_NAME, engine=engine) + filtered = TableRows(rows.table, rows.table.c.v >= 2, engine=engine) + assert len(filtered) == 2 == len(list(filtered))