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
10 changes: 8 additions & 2 deletions sqldol/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
36 changes: 36 additions & 0 deletions sqldol/tests/test_table_rows_len.py
Original file line number Diff line number Diff line change
@@ -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))
Loading