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
48 changes: 43 additions & 5 deletions sqldol/base.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
"""Base objects for sqldol"""

from typing import Union, List
from typing import Union, List, Literal
from collections.abc import Iterable, Iterable, Mapping, Sized, MutableMapping
from sqlalchemy import (
Table,
Expand All @@ -11,6 +11,7 @@
exists,
update,
text,
func,
Engine,
Column,
)
Expand Down Expand Up @@ -167,6 +168,21 @@ def _validate_key_columns(engine, table_name, key_columns):
return key_columns


MissingKeyPolicy = Literal['empty', 'raise']
_missing_key_policies = ('empty', 'raise')


def _validate_missing_key_policy(missing_key_policy) -> MissingKeyPolicy:
"""Validate a ``missing_key_policy`` at construction time, not at lookup time."""
if missing_key_policy not in _missing_key_policies:
msg = (
f'missing_key_policy must be one of {_missing_key_policies}, '
f'not {missing_key_policy!r}'
)
raise ValueError(msg)
return missing_key_policy


# TODO: Make SqlBaseKvReader into a context manager. See https://github.com/i2mint/dol/discussions/49#discussioncomment-8658626
# TODO: Implement filt. (Needs to filt iter, len, and getitem.
# TODO: Refactor idea: Make a query class and a query executor class. Two layers below SqlBaseKvReader.
Expand All @@ -178,6 +194,16 @@ class SqlBaseKvReader(Mapping):
"""A mapping view of a table,
where keys are values from a key column and values are values from a value column.
There's also a filter function that can be used to filter the rows.

The ``missing_key_policy`` keyword decides what a lookup of an absent key does:

- ``'empty'`` (the default) returns an empty result. This is the historical
behavior, kept as the default for backwards compatibility, but note that it
makes the inherited ``__contains__`` and ``get(key, default)`` lie: because
``Mapping`` implements both in terms of ``__getitem__`` raising ``KeyError``,
every key looks present and ``get`` never returns its default.
- ``'raise'`` raises ``KeyError``, which is what ``collections.abc.Mapping``
requires and what makes ``in`` and ``get`` truthful.
"""

def __init__(
Expand All @@ -187,8 +213,11 @@ def __init__(
key_columns: str = None,
value_columns: str | list[str] = None,
filt=None,
*,
missing_key_policy: MissingKeyPolicy = 'empty',
):
self.engine = ensure_engine(engine)
self.missing_key_policy = _validate_missing_key_policy(missing_key_policy)
self.table_name = table_name
key_columns = _validate_key_columns(self.engine, self.table_name, key_columns)
self.metadata = MetaData()
Expand Down Expand Up @@ -222,17 +251,26 @@ def __iter__(self):
yield self._extract_key(row)

def __len__(self):
# Note: A SELECT's ``rowcount`` is -1 on drivers that don't pre-buffer
# results (SQLite, for one), so the row count has to be asked for directly.
query = select(func.count()).select_from(self.table)
with self.engine.connect() as connection:
result = connection.execute(self._table_selection_query)
return result.rowcount
return connection.execute(query).scalar_one()

def __getitem__(self, key):
query = self._table_selection_query.where(self.table.c[self.key_columns] == key)
with self.engine.connect() as connection:
try :
if self.missing_key_policy == 'raise':
rows = connection.execute(query).fetchall()
if not rows:
raise KeyError(key)
return map(self._extract_values, rows)
try:
result = connection.execute(query)
return map(self._extract_values, result.fetchall())
except :
except Exception:
# Note: Narrowed from a bare ``except`` so that KeyboardInterrupt
# and SystemExit are no longer swallowed.
return None

# def __getitem__(self, key):
Expand Down
18 changes: 18 additions & 0 deletions sqldol/sql_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -347,6 +347,24 @@ def __getitem__(self, k):

return doc

def __contains__(self, k) -> bool:
"""True if and only if a row matches the key query ``k``.

Getitem-based, so it agrees with ``__getitem__``, and it costs one
query instead of the brute-force scan inherited from
``dol.base.Collection``. That inherited scan compared ``k`` against
whatever ``__iter__`` yields -- ORM row objects, never keys -- and so
answered False for every key, present or absent.

A key that cannot name a row at all (not a mapping of field to value,
or naming a field the table doesn't have) is simply not contained,
which is what the inherited scan answered too.
"""
try:
return self.query.filter_by(**k).first() is not None
except (TypeError, sqlalchemy.exc.InvalidRequestError):
return False

def __setitem__(self, k, v):
try:
doc = self[k]
Expand Down
2 changes: 1 addition & 1 deletion sqldol/stores.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ def _first_value(iterable):
"""
Get the first element of an iterable.

>>> _get_first([1, 2])
>>> _first_value([1, 2])
1

"""
Expand Down
111 changes: 111 additions & 0 deletions sqldol/tests/test_base_mapping_contract.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,111 @@
"""Tests for the ``Mapping`` contract of ``SqlBaseKvReader`` and its subclasses.

Covers ``__len__`` (which must return a real, non-negative row count) and the
``missing_key_policy`` that decides what an absent key does: return an empty
result (the legacy default) or raise ``KeyError`` (which is what makes the
inherited ``__contains__`` and ``.get(key, default)`` truthful).

Uses in-memory SQLite only -- no external service, no credentials.
"""

import pytest
from sqlalchemy import create_engine, text

from sqldol import SqlDictsReader, SqlDictReader, SqlRowReader, SqlRowsReader

TABLE_NAME = 't'
KEY_COLUMN = 'k'


@pytest.fixture
def engine():
"""An in-memory SQLite engine holding ``t(k, v)`` with rows a/1 and b/2."""
engine = create_engine('sqlite:///:memory:')
with engine.connect() as connection:
connection.execute(text(f'CREATE TABLE {TABLE_NAME} (k TEXT, v TEXT)'))
connection.execute(text(f"INSERT INTO {TABLE_NAME} VALUES ('a', '1'), ('b', '2')"))
connection.commit()
return engine


@pytest.mark.parametrize(
'cls', [SqlRowsReader, SqlRowReader, SqlDictsReader, SqlDictReader]
)
def test_len_is_a_real_row_count(engine, cls):
"""``__len__`` must count rows, not return a SELECT's ``rowcount``.

Before the fix every one of these raised
``ValueError: __len__() should return >= 0``, because a SELECT reports
``rowcount == -1`` on SQLite (and any other driver that does not
pre-buffer results).
"""
store = cls(engine, TABLE_NAME, key_columns=KEY_COLUMN)
assert len(store) == 2


@pytest.mark.parametrize(
'cls', [SqlRowsReader, SqlRowReader, SqlDictsReader, SqlDictReader]
)
def test_list_of_store_yields_the_keys(engine, cls):
"""``list(store)`` used to die on the ``__len__`` length hint, though
``iter(store)`` worked fine."""
store = cls(engine, TABLE_NAME, key_columns=KEY_COLUMN)
assert sorted(store) == ['a', 'b']
assert sorted(list(store)) == ['a', 'b']


def test_values_of_present_keys_are_unchanged(engine):
"""Guard that nothing about the happy path moved."""
assert SqlDictsReader(engine, TABLE_NAME, key_columns=KEY_COLUMN)['a'] == [
{'k': 'a', 'v': '1'}
]
assert SqlRowsReader(engine, TABLE_NAME, key_columns=KEY_COLUMN)['a'] == [
['a', '1']
]
assert SqlRowReader(engine, TABLE_NAME, key_columns=KEY_COLUMN)['a'] == ['a', '1']
assert SqlDictReader(engine, TABLE_NAME, key_columns=KEY_COLUMN)['a'] == {
'k': 'a',
'v': '1',
}


@pytest.mark.parametrize(
'cls', [SqlRowsReader, SqlRowReader, SqlDictsReader, SqlDictReader]
)
def test_missing_key_policy_raise_makes_the_mapping_contract_truthful(engine, cls):
"""With ``missing_key_policy='raise'``, ``in`` and ``.get`` stop lying.

Before the fix ``'zzz' in store`` was ``True`` for every store (because
``Mapping.__contains__`` is getitem-based and getitem never raised), and
``.get('zzz', default)`` returned an empty result instead of ``default``.
"""
store = cls(engine, TABLE_NAME, key_columns=KEY_COLUMN, missing_key_policy='raise')

with pytest.raises(KeyError):
store['zzz']

assert 'zzz' not in store
assert 'a' in store
assert store.get('zzz', 'DEFAULT') == 'DEFAULT'
assert store.get('a') is not None


def test_missing_key_policy_defaults_to_the_legacy_empty_behavior(engine):
"""Regression guard: the default must keep 2024-era callers working.

Existing consumers branch on ``SqlDictReader[missing] is None``, so the
``KeyError`` behaviour has to stay opt-in.
"""
assert SqlDictReader(engine, TABLE_NAME, key_columns=KEY_COLUMN)['zzz'] is None
assert SqlDictsReader(engine, TABLE_NAME, key_columns=KEY_COLUMN)['zzz'] == []
assert SqlRowsReader(engine, TABLE_NAME, key_columns=KEY_COLUMN)['zzz'] == []
with pytest.raises(StopIteration):
SqlRowReader(engine, TABLE_NAME, key_columns=KEY_COLUMN)['zzz']


def test_unknown_missing_key_policy_is_rejected_at_construction(engine):
"""A typo must fail loudly, and at construction time, not at lookup time."""
with pytest.raises(ValueError):
SqlDictsReader(
engine, TABLE_NAME, key_columns=KEY_COLUMN, missing_key_policy='nope'
)
83 changes: 83 additions & 0 deletions sqldol/tests/test_sqlalchemy_store_contract.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,83 @@
"""Tests that ``k in store`` agrees with ``store[k]`` for the SQLAlchemy stores.

``SQLAlchemyPersister.__iter__`` yields ORM row objects rather than keys, so the
brute-force ``__contains__`` inherited from ``dol.base.Collection`` compared a key
against a row and answered ``False`` for every key. ``SQLAlchemyStore`` and
``SQLAlchemyTupleStore`` inherited that answer through ``dol.base.Store``.

Uses in-memory SQLite only -- no external service, no credentials.
"""

import pytest

from sqldol import SQLAlchemyPersister, SQLAlchemyStore, SQLAlchemyTupleStore

SQLITE_DB_URI = 'sqlite:///:memory:'
KEY_FIELDS = {'id': SQLAlchemyPersister.TYPE_STRING}
DATA_FIELDS = {'doc': SQLAlchemyPersister.TYPE_TEXT}


def _mk(cls, collection_name):
return cls(
uri=SQLITE_DB_URI,
collection_name=collection_name,
key_fields=KEY_FIELDS,
data_fields=DATA_FIELDS,
)


def test_persister_contains_agrees_with_getitem():
persister = _mk(SQLAlchemyPersister, 'persister_contains')
persister[{'id': 'a'}] = {'doc': 'x'}

assert {'id': 'a'} in persister
assert {'id': 'zzz'} not in persister
assert persister[{'id': 'a'}] is not None
assert len(persister) == 1


def test_store_contains_agrees_with_getitem():
store = _mk(SQLAlchemyStore, 'store_contains')
store[{'id': 'a'}] = {'doc': 'x'}

assert {'id': 'a'} in store
assert {'id': 'zzz'} not in store
assert store[{'id': 'a'}] is not None
assert len(store) == 1


def test_tuple_store_contains_agrees_with_getitem():
store = _mk(SQLAlchemyTupleStore, 'tuple_store_contains')
store[('a',)] = ('x',)

assert ('a',) in store
assert ('zzz',) not in store
assert store[('a',)] == ('x',)
assert len(store) == 1


def test_tuple_store_iteration_still_yields_keys():
"""Pins that ``__iter__`` was deliberately left alone.

Making the persister's ``__iter__`` yield keys instead of ORM rows is a
separate, not-backwards-compatible change: ``SQLAlchemyTupleStore._key_of_id``
does ``getattr(obj, field)`` on whatever is yielded, and at least one known
consumer wraps the row-yielding behaviour.
"""
store = _mk(SQLAlchemyTupleStore, 'tuple_store_iter')
store[('a',)] = ('x',)

assert list(store) == [('a',)]


def test_contains_does_not_raise_on_a_key_it_cannot_query():
"""``in`` must answer, not explode, for keys that can't name a row.

``Container.__contains__`` returned False for these before, so keeping them
False (rather than propagating a SQLAlchemy/TypeError) preserves behaviour.
"""
persister = _mk(SQLAlchemyPersister, 'persister_bad_key')
persister[{'id': 'a'}] = {'doc': 'x'}

assert 'not-a-dict' not in persister
assert {'no_such_column': 'a'} not in persister
Loading