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
58 changes: 45 additions & 13 deletions sqldol/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,10 +10,10 @@
insert,
exists,
update,
text,
func,
Engine,
Column,
and_,
)

from sqlalchemy import Table, Column, MetaData
Expand Down Expand Up @@ -291,19 +291,51 @@ def __getitem__(self, key):
# TODO: Needs to be made compliant with the "Base" strategy (see SqlBaseKvReader)
# For example, perhaps values are not dicts, but lists of rows
class SqlBaseKvStore(SqlBaseKvReader, MutableMapping):
def _key_column(self, column_name):
"""The ``Column`` object named ``column_name``, or an informative ``ValueError``.

Going through ``self.table.c`` means a column name is never spliced into SQL
text: only names the table actually has can be used, and SQLAlchemy quotes
them as the dialect requires. As with unquoted names in SQL, a name that
matches no column exactly but exactly one column case-insensitively
resolves to that column.

An unknown column is a ``ValueError``, not a ``KeyError``, so that a typo is
not mistaken for "that key is absent" by callers catching ``KeyError``.
"""
if column_name in self.table.c:
return self.table.c[column_name]
if isinstance(column_name, str):
matches = [c for c in self.table.c if c.name.lower() == column_name.lower()]
if len(matches) == 1:
return matches[0]
msg = (
f"{column_name!r} is not a column of table {self.table_name!r}. "
f"Its columns are {self._column_names}."
)
raise ValueError(msg)

def _mk_column_filter(self, key):
if isinstance(key, str):
return text(f"{self.key_columns} = '{key}'")
elif isinstance(key, int): # the key is a tuple of columns
return text(f"{self.key_columns} = {key}")
elif isinstance(key, dict):
return text(" AND ".join(f"{col} = '{val}'" for col, val in key.items()))
else:
return text(
" AND ".join(
f"{col} = '{val}'" for col, val in zip(self.key_columns, key)
)
"""The ``WHERE`` clause selecting the rows of ``key``.

A ``Mapping`` key is a ``{column_name: value, ...}`` conjunction; any other key
is the value of the (single) key column. Values are always bound parameters,
never interpolated into SQL text, so keys containing quotes, semicolons or
``--`` are matched literally. For non-mapping keys this is the same
comparison ``__getitem__`` makes.
"""
if isinstance(key, Mapping):
if not key:
raise ValueError("A mapping key must name at least one column")
return and_(*(self._key_column(col) == val for col, val in key.items()))
if isinstance(key, (tuple, list)):
msg = (
f"Only a single key column ({self.key_columns!r}) is supported, so a "
f"key can't be a {type(key).__name__}. To match several columns, use "
f"a mapping key: {{column_name: value, ...}}."
)
raise TypeError(msg)
return self._key_column(self.key_columns) == key

def __setitem__(self, key, value):
filter = self._mk_column_filter(key)
Expand All @@ -327,7 +359,7 @@ def __setitem__(self, key, value):

# TODO: Should we return something useful?

def __delitem__(self, key: str) -> None:
def __delitem__(self, key) -> None:
filter = self._mk_column_filter(key)
query = delete(self.table).where(filter)

Expand Down
48 changes: 47 additions & 1 deletion sqldol/sql_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,9 @@
sql with a simple (dict-like or list-like) interface
"""

import re
from functools import partial
from operator import index

import sqlalchemy
from sqlalchemy import create_engine, Column, String, Table
Expand All @@ -18,6 +20,43 @@
DFLT_SQL_PORT = 1433
DFLT_SQL_HOST = "localhost"

_SQL_IDENTIFIER_PART = r"(?=[\w$]*[^\W\d])[\w$]+" # word chars or $, not all digits
SQL_IDENTIFIER_PATTERN = re.compile(
rf"{_SQL_IDENTIFIER_PART}(\.{_SQL_IDENTIFIER_PART})?"
)
"""What a table name may look like where it is written into raw SQL text: letters
(Unicode included), digits, ``_`` or ``$``, not all digits (MySQL allows e.g.
``2020_sales``), optionally qualified by a schema (``schema.table``). None of these
characters can end an identifier or start a new SQL token."""


def validate_sql_identifier(name, *, pattern=SQL_IDENTIFIER_PATTERN):
"""Return ``name`` if it is safe to write into raw SQL text, else raise ``ValueError``.

The raw-SQL paths of this module cannot bind a table name as a parameter (no SQL
dialect allows that), and they may be handed a plain DB-API connection that has no
quoting helper, so they only accept names matching an allowlist.

>>> validate_sql_identifier("my_table")
'my_table'
>>> validate_sql_identifier("my_schema.my_table")
'my_schema.my_table'
>>> validate_sql_identifier("t; SELECT 1") # doctest: +ELLIPSIS
Traceback (most recent call last):
...
ValueError: Not a valid SQL table name: 't; SELECT 1'. ...
"""
if not isinstance(name, str) or not pattern.fullmatch(name):
msg = (
f"Not a valid SQL table name: {name!r}. "
f"Expected letters, digits, '_' or '$' (not all digits), "
f"optionally qualified as 'schema.table'. "
f"Names that need quoting are not supported here: use the SQLAlchemy-based "
f"stores (sqldol.base, sqldol.stores), which quote identifiers themselves."
)
raise ValueError(msg)
return name


# TODO: decorator to automatically retry (once) if the connection times out
# --> from sqlalchemy.exc import OperationalError
Expand All @@ -44,7 +83,9 @@ class SqlTableRowsCollection(Collection):

def __init__(self, connection, table_name, batch_size=2000, limit=int(1e16)):
self.connection = connection
self.table_name = table_name
# Note: table_name is written into raw SQL text (see the _tmpl_* templates),
# so it has to pass the identifier allowlist.
self.table_name = validate_sql_identifier(table_name)
self.iter_rows = partial(
iter_rows,
connection=connection,
Expand Down Expand Up @@ -398,7 +439,12 @@ def __len__(self):
def iter_rows(connection, table_name, batch_size=1000, offset=0, limit=int(1e12)):
"""Iterate the over the rows of a table.
The limit argument is mostly there to avoid an infinite loop, but can also be used to get ranges.

``table_name`` must pass :func:`validate_sql_identifier`, and ``batch_size``,
``offset`` and ``limit`` must be integers, because they are written into the SQL text.
"""
table_name = validate_sql_identifier(table_name)
batch_size, offset, limit = index(batch_size), index(offset), index(limit)
stop_offset = limit + offset # to the range act like like sql limit
i = 0
for offset in range(offset, stop_offset, batch_size):
Expand Down
Loading
Loading