diff --git a/lnbits/db.py b/lnbits/db.py index 1a6c562ee..55c98ab9a 100644 --- a/lnbits/db.py +++ b/lnbits/db.py @@ -35,7 +35,7 @@ if settings.lnbits_database_url: else: if not database_uri.startswith("postgres://"): raise ValueError( - "Please use the 'postgres://...' " "format for the database URL." + "Please use the 'postgres://...' format for the database URL." ) DB_TYPE = POSTGRES @@ -560,7 +560,9 @@ class Filters(BaseModel, Generic[TFilterModel]): def pagination(self) -> str: stmt = "" - self.limit = self.limit or 10 + if self.limit == 0: + self.limit = 1000 + self.limit = 10 if self.limit is None else self.limit stmt += f"LIMIT {min(1000, self.limit)} " if self.offset: stmt += f"OFFSET {self.offset}" diff --git a/tests/unit/test_db_fetch_page.py b/tests/unit/test_db_fetch_page.py index c110332e4..6cca11fa7 100644 --- a/tests/unit/test_db_fetch_page.py +++ b/tests/unit/test_db_fetch_page.py @@ -1,7 +1,32 @@ import pytest +from lnbits.db import Filters from tests.helpers import DbTestModel +TEST_DB_FETCH_PAGE_ROWS: tuple[dict[str, str], ...] = ( + {"id": "1", "name": "Alice", "value": "foo"}, + {"id": "2", "name": "Bob", "value": "bar"}, + {"id": "3", "name": "Carol", "value": "bar"}, + {"id": "4", "name": "Dave", "value": "bar"}, + {"id": "5", "name": "Dave", "value": "foo"}, + {"id": "6", "name": "Eve", "value": "foo"}, + {"id": "7", "name": "Frank", "value": "bar"}, + {"id": "8", "name": "Grace", "value": "foo"}, + {"id": "9", "name": "Heidi", "value": "bar"}, + {"id": "10", "name": "Ivan", "value": "foo"}, + {"id": "11", "name": "Judy", "value": "bar"}, + {"id": "12", "name": "Mallory", "value": "foo"}, + {"id": "13", "name": "Niaj", "value": "bar"}, + {"id": "14", "name": "Olivia", "value": "foo"}, + {"id": "15", "name": "Peggy", "value": "bar"}, + {"id": "16", "name": "Rupert", "value": "foo"}, + {"id": "17", "name": "Sybil", "value": "bar"}, + {"id": "18", "name": "Trent", "value": "foo"}, + {"id": "19", "name": "Victor", "value": "bar"}, + {"id": "20", "name": "Walter", "value": "foo"}, + {"id": "21", "name": "Zoe", "value": "bar"}, +) + @pytest.fixture(scope="session") async def fetch_page(db): @@ -13,14 +38,14 @@ async def fetch_page(db): name TEXT NOT NULL ) """) - await db.execute(""" - INSERT INTO test_db_fetch_page (id, name, value) VALUES - ('1', 'Alice', 'foo'), - ('2', 'Bob', 'bar'), - ('3', 'Carol', 'bar'), - ('4', 'Dave', 'bar'), - ('5', 'Dave', 'foo') - """) + for row in TEST_DB_FETCH_PAGE_ROWS: + await db.execute( + """ + INSERT INTO test_db_fetch_page (id, name, value) + VALUES (:id, :name, :value) + """, + row, + ) yield await db.execute("DROP TABLE test_db_fetch_page") @@ -33,8 +58,35 @@ async def test_db_fetch_page_simple(fetch_page, db): ) assert row - assert row.total == 5 - assert len(row.data) == 5 + assert row.total == len(TEST_DB_FETCH_PAGE_ROWS) + assert len(row.data) == Filters().limit + + +@pytest.mark.anyio +async def test_db_fetch_page_limit_zero_returns_all(fetch_page, db): + row = await db.fetch_page( + query="select * from test_db_fetch_page", + filters=Filters(limit=0), + model=DbTestModel, + ) + + assert row + assert row.total == len(TEST_DB_FETCH_PAGE_ROWS) + assert len(row.data) == len(TEST_DB_FETCH_PAGE_ROWS) + + +@pytest.mark.anyio +async def test_db_fetch_page_limit(fetch_page, db): + limit = 5 + row = await db.fetch_page( + query="select * from test_db_fetch_page", + filters=Filters(limit=limit), + model=DbTestModel, + ) + + assert row + assert row.total == len(TEST_DB_FETCH_PAGE_ROWS) + assert len(row.data) == limit @pytest.mark.anyio @@ -45,7 +97,7 @@ async def test_db_fetch_page_group_by(fetch_page, db): group_by=["name"], ) assert row - assert row.total == 4 + assert row.total == len({test_row["name"] for test_row in TEST_DB_FETCH_PAGE_ROWS}) @pytest.mark.anyio @@ -56,7 +108,9 @@ async def test_db_fetch_page_group_by_multiple(fetch_page, db): group_by=["value", "name"], ) assert row - assert row.total == 5 + assert row.total == len( + {(test_row["value"], test_row["name"]) for test_row in TEST_DB_FETCH_PAGE_ROWS} + ) @pytest.mark.anyio