Skip to content
Open
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
15 changes: 15 additions & 0 deletions .github/workflows/main_checks.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,21 @@ jobs:
build:
runs-on: ubuntu-latest

services:
postgres:
image: postgres:16
env:
POSTGRES_USER: root
POSTGRES_PASSWORD: root
POSTGRES_DB: lematerial
ports:
- 5432:5432
options: >-
--health-cmd pg_isready
--health-interval 10s
--health-timeout 5s
--health-retries 5

steps:
- name: Checkout code
uses: actions/checkout@v3
Expand Down
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,7 @@ material-hasher = { git = "https://github.com/LeMaterial/lematerial-hasher.git"
[tool.pytest.ini_options]
markers = [
"integration: tests that require real AWS credentials and S3 access (deselect with '-m \"not integration\"')",
"postgres: tests that require a running PostgreSQL server, skipped automatically if none is reachable (deselect with '-m \"not postgres\"')",
]

[tool.ruff.lint]
Expand Down
6 changes: 4 additions & 2 deletions src/lematerial_fetcher/database/postgres.py
Original file line number Diff line number Diff line change
Expand Up @@ -285,12 +285,14 @@ def fetch_items_iter(
if not start_id: # No results at this offset
return

# Construct the query based on whether we have a start_id
# Construct the query based on whether we have a start_id.
# start_id is the id *at* the requested offset (0-based), so the
# scan must be inclusive: `>` would drop the boundary row itself.
if start_id:
query = f"""
SELECT id, type, attributes, last_modified
FROM {table_name}
WHERE id > %s
WHERE id >= %s
ORDER BY id
{f"LIMIT {limit}" if limit is not None else ""}
"""
Expand Down
125 changes: 125 additions & 0 deletions tests/database/test_postgres.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,125 @@
# Copyright 2025 Entalpic
"""Regression tests for offset pagination in ``StructuresDatabase``.

These tests run against a real PostgreSQL server, because the bug they guard
against lives in the SQL itself (``WHERE id > start_id`` vs ``WHERE id >= start_id``)
and a mocked cursor cannot catch it.

They are skipped automatically when no server is reachable. To run them locally:

.. code-block:: bash

docker run -d --name lemat-test-pg \\
-e POSTGRES_USER=root -e POSTGRES_PASSWORD=root -e POSTGRES_DB=lematerial \\
-p 5432:5432 postgres:16
uv run pytest tests/database/ -m postgres

Set ``LEMATERIALFETCHER_TEST_DB_CONN_STR`` to point at a different server.
"""

import os
import uuid

import psycopg2
import pytest

from lematerial_fetcher.database.postgres import StructuresDatabase
from lematerial_fetcher.models.models import RawStructure

DEFAULT_CONN_STR = (
"host=localhost user=root password=root dbname=lematerial sslmode=disable"
)
CONN_STR = os.environ.get("LEMATERIALFETCHER_TEST_DB_CONN_STR", DEFAULT_CONN_STR)

pytestmark = pytest.mark.postgres


def _postgres_available() -> bool:
try:
psycopg2.connect(CONN_STR).close()
return True
except psycopg2.OperationalError:
return False


@pytest.fixture
def db():
"""A StructuresDatabase backed by a uniquely-named, throwaway table."""
if not _postgres_available():
pytest.skip(
"No PostgreSQL server available "
"(set LEMATERIALFETCHER_TEST_DB_CONN_STR or start a local server)"
)
table_name = f"test_structures_{uuid.uuid4().hex[:12]}"
database = StructuresDatabase(CONN_STR, table_name)
database.create_table()
yield database
with database.conn.cursor() as cur:
cur.execute(f"DROP TABLE IF EXISTS {table_name};")
database.conn.commit()
database.close()


def insert_rows(db: StructuresDatabase, n: int) -> list[str]:
"""Insert n rows with ids that sort in insertion order, return the ids."""
ids = [f"struct-{i:03d}" for i in range(n)]
db.batch_insert_data(
[
RawStructure(
id=id_,
type="test-structure",
attributes={"index": i},
last_modified=None,
)
for i, id_ in enumerate(ids)
]
)
return ids


def test_fetch_items_iter_starts_at_offset(db):
"""fetch_items_iter(offset=k) must start at the row at rank k, not k+1."""
ids = insert_rows(db, 10)
got = [row.id for row in db.fetch_items_iter(offset=4, limit=3)]
assert got == ids[4:7]


def test_fetch_items_returns_row_at_offset(db):
ids = insert_rows(db, 10)
rows = db.fetch_items(offset=3, batch_size=1)
assert [row.id for row in rows] == [ids[3]]


def test_batched_pagination_covers_every_row_exactly_once(db):
"""Reading a table in consecutive (offset, batch_size) windows, the way
BaseTransformer does, must yield every row exactly once.

Regression test: fetch_items_iter used to resolve the id *at* the offset
and then scan ``WHERE id > start_id``, silently dropping the row at rank
``batch_size`` from every multi-batch read.
"""
ids = insert_rows(db, 10)
batch_size = 3

seen = []
offset = 0
while True:
batch = db.fetch_items(offset=offset, batch_size=batch_size)
if not batch:
break
seen.extend(row.id for row in batch)
offset += batch_size

assert seen == ids


def test_offset_zero_reads_from_first_row(db):
ids = insert_rows(db, 5)
got = [row.id for row in db.fetch_items_iter(offset=0)]
assert got == ids


def test_offset_past_end_yields_nothing(db):
insert_rows(db, 5)
assert list(db.fetch_items_iter(offset=5)) == []
assert db.fetch_items(offset=99, batch_size=10) == []