diff --git a/odbcdol/__init__.py b/odbcdol/__init__.py index 73cf6ce..c7f3320 100644 --- a/odbcdol/__init__.py +++ b/odbcdol/__init__.py @@ -3,6 +3,7 @@ """ import platform +import re import subprocess from collections.abc import MutableMapping @@ -12,6 +13,65 @@ pyodbc = check_pyodbc_import() +_MAX_IDENTIFIER_LEN = 128 + + +def _quote_identifier(name: str) -> str: + """Return *name* as a bracket-quoted SQL Server identifier. + + Identifiers cannot be passed as query parameters, so they are delimited: + ``]`` is doubled, which makes any character sequence a single identifier. + A name the caller already bracketed (``[first name]``) is unwrapped first. + Empty names, names over 128 characters and control characters are refused. + """ + if not isinstance(name, str): + raise ValueError(f"Not a valid SQL identifier: {name!r}") + if len(name) > 2 and name.startswith("[") and name.endswith("]"): + name = name[1:-1].replace("]]", "]") + if ( + not name + or len(name) > _MAX_IDENTIFIER_LEN + or any(ord(c) < 0x20 or ord(c) == 0x7F for c in name) + ): + raise ValueError(f"Not a valid SQL identifier: {name!r}") + return "[" + name.replace("]", "]]") + "]" + + +def _quote_table_name(name: str) -> str: + """Quote a possibly schema-qualified table name (``dbo.person``), part by part.""" + if not isinstance(name, str): + raise ValueError(f"Not a valid SQL table name: {name!r}") + parts = name.split(".") + if not 1 <= len(parts) <= 4: + raise ValueError(f"Not a valid SQL table name: {name!r}") + return ".".join(_quote_identifier(p) for p in parts) + + +_SCALAR_TYPES = (str, int, float, bool, bytes, bytearray, type(None)) + + +def _bindable(value): + """A value pyodbc can bind as one parameter. + + Scalars (and dates, decimals...) pass through; containers are stored as + their ``str()``, as before. (A lone list/tuple would otherwise be read by + pyodbc as the whole parameter list.) + """ + if isinstance(value, (list, tuple, dict, set, frozenset)): + return str(value) + return value + + +def _connection_string(**parts) -> str: + """Build an ODBC connection string with every value brace-quoted. + + Braced values may contain ``;`` and ``=`` (``}`` is doubled), so no value + -- a password such as ``x};Encrypt=no;APP={y`` included -- can end its + attribute and start another. + """ + return ";".join(f"{k}={{{str(v).replace('}', '}}')}}}" for k, v in parts.items()) + + class SQLServerPersister(MutableMapping): def __init__( self, @@ -28,31 +88,28 @@ def __init__( self.__check_dependencies() self._sql_server_client = pyodbc.connect( - "DRIVER={{ODBC Driver 17 for SQL Server}};" - "SERVER={}:{},{};" - "DATABASE={};" - "UID={};" - "PWD={}".format(conn_protocol, host, port, db_name, db_username, db_pass) + _connection_string( + DRIVER="ODBC Driver 17 for SQL Server", + SERVER=f"{conn_protocol}:{host},{port}", + DATABASE=db_name, + UID=db_username, + PWD=db_pass, + ) ) self._cursor = self._sql_server_client.cursor() self._table_name = table_name self._primary_key = primary_key - self._select_all_query = "SELECT * from {table};".format(table=self._table_name) - self._insert_query = ( - "INSERT into {table}({{attributes}}) VALUES ({{values}});".format( - table=self._table_name - ) - ) - self._select_query = ( - "SELECT * from {table} where {primary_key}='{{value}}'".format( - table=self._table_name, primary_key=self._primary_key - ) - ) - self._del_query = "DELETE from {table} where {primary_key} = {{value}};".format( - table=self._table_name, primary_key=self._primary_key - ) + # Identifiers cannot be bound as parameters, so they are bracket-quoted + # (which also keeps braces out of the templates below: the insert + # template is .format-ted later); every VALUE is bound (``?``). + table = _quote_table_name(self._table_name) + pk = _quote_identifier(self._primary_key) + self._select_all_query = f"SELECT * FROM {table};" + self._insert_query = f"INSERT INTO {table}({{attributes}}) VALUES ({{values}});" + self._select_query = f"SELECT * FROM {table} WHERE {pk} = ?;" + self._del_query = f"DELETE FROM {table} WHERE {pk} = ?;" @staticmethod def __check_dependencies(): @@ -82,21 +139,20 @@ def __check_dependencies(): ) def __getitem__(self, k): - self._cursor.execute(self._select_query.format(value=k)) + self._cursor.execute(self._select_query, (_bindable(k),)) record = self._cursor.fetchone() return record if record else print(f"No record found for primary_key: {k}") # TODO: Raise a proper exception here def __setitem__(self, k, v): - _sanitized_values = [ - str(val) if isinstance(val, int) else f"'{val}'" for val in v.values() - ] + columns = list(v.keys()) try: self._cursor.execute( self._insert_query.format( - attributes=",".join(v.keys()), - values=",".join(_sanitized_values), - ) + attributes=",".join(_quote_identifier(c) for c in columns), + values=",".join("?" for _ in columns), + ), + tuple(_bindable(v[c]) for c in columns), ) except pyodbc.IntegrityError as e: # TODO: Raise a proper exception here @@ -105,7 +161,7 @@ def __setitem__(self, k, v): self._sql_server_client.commit() def __delitem__(self, k): - self._cursor.execute(self._del_query.format(value=k)) + self._cursor.execute(self._del_query, (_bindable(k),)) self._sql_server_client.commit() def __iter__(self): diff --git a/odbcdol/tests/test_injection.py b/odbcdol/tests/test_injection.py new file mode 100644 index 0000000..6301a1e --- /dev/null +++ b/odbcdol/tests/test_injection.py @@ -0,0 +1,148 @@ +"""Values reach SQL Server as bound parameters, never as SQL text. + +Runs without a database: a stand-in ``pyodbc`` records what would be sent. +""" + +import sys +import types + +import pytest + + +class _Cursor: + def __init__(self, log): + self.log = log + + def execute(self, sql, params=()): + self.log.append((sql, tuple(params))) + + def fetchone(self): + return ("row",) + + def fetchall(self): + return [] + + +class _Conn: + def __init__(self, conn_str, log): + self.conn_str = conn_str + self.log = log + + def cursor(self): + return _Cursor(self.log) + + def commit(self): + pass + + +@pytest.fixture +def persister_cls(monkeypatch): + log = [] + fake = types.SimpleNamespace( + connect=lambda s: _Conn(s, log), IntegrityError=type("IE", (Exception,), {}) + ) + monkeypatch.setitem(sys.modules, "pyodbc", fake) + import odbcdol + + monkeypatch.setattr(odbcdol, "pyodbc", fake) + monkeypatch.setattr( + odbcdol.SQLServerPersister, + "_SQLServerPersister__check_dependencies", + staticmethod(lambda: None), + ) + return odbcdol.SQLServerPersister, log + + +HOSTILE = ["1 OR 1=1", "x'; DROP TABLE person; --", "1; DELETE FROM person"] + + +@pytest.mark.parametrize("key", HOSTILE) +def test_keys_are_bound_on_read_and_delete(persister_cls, key): + cls, log = persister_cls + s = cls() + s[key] + del s[key] + for sql, params in log: + assert key not in sql + assert params == (key,) + assert log[0][0] == "SELECT * FROM [person] WHERE [id] = ?;" + assert log[1][0] == "DELETE FROM [person] WHERE [id] = ?;" + + +@pytest.mark.parametrize("value", HOSTILE) +def test_values_are_bound_on_write(persister_cls, value): + cls, log = persister_cls + s = cls() + s["3"] = {"id": "3", "name": value} + sql, params = log[-1] + assert sql == "INSERT INTO [person]([id],[name]) VALUES (?,?);" + assert params == ("3", value) + + +@pytest.mark.parametrize( + "column, quoted", + [ + ("name) VALUES (1); --", "[name) VALUES (1); --]"), + ("x]; DROP TABLE t; --", "[x]]; DROP TABLE t; --]"), + ("first name", "[first name]"), + ("[first name]", "[first name]"), + ], +) +def test_column_names_are_always_one_identifier(persister_cls, column, quoted): + cls, log = persister_cls + cls()["3"] = {column: "v"} + sql, params = log[-1] + assert sql == f"INSERT INTO [person]({quoted}) VALUES (?);" + assert params == ("v",) + + +@pytest.mark.parametrize("column", ["", "a\x00b", "x" * 129, 3]) +def test_invalid_column_names_are_refused(persister_cls, column): + cls, _ = persister_cls + with pytest.raises(ValueError): + cls()["3"] = {column: "v"} + + +@pytest.mark.parametrize( + "table, quoted", + [ + ("person", "[person]"), + ("dbo.person", "[dbo].[person]"), + ("[dbo].[person]", "[dbo].[person]"), + ("person]; DROP TABLE x; --", "[person]]; DROP TABLE x; --]"), + ], +) +def test_table_names_quote_part_by_part(persister_cls, table, quoted): + cls, log = persister_cls + s = cls(table_name=table) + s["k"] + assert log[-1][0] == f"SELECT * FROM {quoted} WHERE [id] = ?;" + + +@pytest.mark.parametrize("key", [(1, 2), [1], {"a": 1}]) +def test_container_keys_bind_as_one_value(persister_cls, key): + cls, log = persister_cls + cls()[key] + assert log[-1][1] == (str(key),) + + +@pytest.mark.parametrize( + "password, quoted", + [ + ("pw;DATABASE=master", "PWD={pw;DATABASE=master}"), + ("{x};Encrypt=no;APP={y}", "PWD={{x}};Encrypt=no;APP={y}}}"), + ("{literalbraces}", "PWD={{literalbraces}}}"), + ], +) +def test_connection_string_values_cannot_add_attributes( + persister_cls, password, quoted +): + cls, _ = persister_cls + s = cls(db_pass=password, db_name="db") + conn_str = s._sql_server_client.conn_str + assert conn_str.endswith(quoted) + assert "DATABASE={db};" in conn_str + assert conn_str.startswith("DRIVER={ODBC Driver 17 for SQL Server};") + assert conn_str.count("Encrypt=no") <= 1 and ";Encrypt=no;" not in conn_str.replace( + quoted, "" + )