From bf2f04849521a23221290425feb9fcc00967eb21 Mon Sep 17 00:00:00 2001 From: Thor Whalen <1906276+thorwhalen@users.noreply.github.com> Date: Tue, 22 Sep 2026 15:11:02 +0000 Subject: [PATCH 1/2] Bind values as query parameters; validate identifiers Keys and values were formatted into SQL text (the delete did not even quote the key), and connection-string values were not brace-quoted. Values are now bound with ?, table/column names are validated and bracket-quoted, and connection-string values that could start a new attribute are brace-quoted. Co-Authored-By: Claude Opus 5 --- odbcdol/__init__.py | 83 ++++++++++++++++--------- odbcdol/tests/test_injection.py | 104 ++++++++++++++++++++++++++++++++ 2 files changed, 160 insertions(+), 27 deletions(-) create mode 100644 odbcdol/tests/test_injection.py diff --git a/odbcdol/__init__.py b/odbcdol/__init__.py index 73cf6ce..79fe6ab 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,39 @@ pyodbc = check_pyodbc_import() +_IDENTIFIER_RE = re.compile(r"[A-Za-z_@#][A-Za-z0-9_@#$]{0,127}") + + +def _quote_identifier(name: str) -> str: + """Validate a table/column name and return it bracket-quoted for SQL Server. + + Identifiers cannot be passed as query parameters, so anything that is not a + plain (regular) SQL Server identifier is refused rather than escaped. + """ + if not isinstance(name, str) or not _IDENTIFIER_RE.fullmatch(name): + raise ValueError(f"Not a valid SQL identifier: {name!r}") + return f"[{name}]" + + +def _connection_string(**parts) -> str: + """Build an ODBC connection string, brace-quoting values that need it. + + A value containing ``;``, ``{``, ``}``, ``=`` or surrounding spaces would + otherwise end its attribute and start another (e.g. a password + ``x;DATABASE=other``). Values already wrapped in braces are kept as is. + """ + + def _value(v) -> str: + v = str(v) + if v.startswith("{") and v.endswith("}"): + return v + if any(c in v for c in ";{}=") or v != v.strip(): + return "{" + v.replace("}", "}}") + "}" + return v + + return ";".join(f"{k}={_value(v)}" for k, v in parts.items()) + + class SQLServerPersister(MutableMapping): def __init__( self, @@ -28,31 +62,27 @@ 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 validated and + # bracket-quoted; every VALUE is bound (``?``) and never formatted in. + table = _quote_identifier(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 +112,20 @@ def __check_dependencies(): ) def __getitem__(self, k): - self._cursor.execute(self._select_query.format(value=k)) + self._cursor.execute(self._select_query, 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), + ), + *(v[c] for c in columns), ) except pyodbc.IntegrityError as e: # TODO: Raise a proper exception here @@ -105,7 +134,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, 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..118192e --- /dev/null +++ b/odbcdol/tests/test_injection.py @@ -0,0 +1,104 @@ +"""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, 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", ["name) VALUES (1); --", "a b", "x]", ""]) +def test_hostile_column_names_are_refused(persister_cls, column): + cls, _ = persister_cls + with pytest.raises(ValueError): + cls()["3"] = {column: "v"} + + +@pytest.mark.parametrize( + "kw", [{"table_name": "person; DROP TABLE x"}, {"primary_key": "id = id OR 1"}] +) +def test_hostile_identifiers_are_refused_at_construction(persister_cls, kw): + cls, _ = persister_cls + with pytest.raises(ValueError): + cls(**kw) + + +def test_connection_string_values_cannot_add_attributes(persister_cls): + cls, log = persister_cls + s = cls(db_pass="pw;DATABASE=master", db_name="db") + conn_str = s._sql_server_client.conn_str + assert "PWD={pw;DATABASE=master}" in conn_str + assert "DATABASE=db;" in conn_str + assert conn_str.startswith("DRIVER={ODBC Driver 17 for SQL Server};") From 1b7fdc8c7b398c5503248643909d0cfd814ae542 Mon Sep 17 00:00:00 2001 From: Thor Whalen <1906276+thorwhalen@users.noreply.github.com> Date: Tue, 22 Sep 2026 15:14:01 +0000 Subject: [PATCH 2/2] Address review: always brace-quote, schema-qualified tables - Every connection-string value is brace-quoted (a value already wrapped in braces could still inject attributes); DRIVER is passed bare. - Identifiers are delimited ([...] with ] doubled) instead of pattern- checked, so any name is one identifier; table names may be schema- qualified (dbo.person) and pre-bracketed names are accepted. - Parameters are passed as one tuple; container values/keys bind as their str() as before, never as pyodbc's parameter list. Co-Authored-By: Claude Opus 5 --- odbcdol/__init__.py | 79 ++++++++++++++++++++++----------- odbcdol/tests/test_injection.py | 70 +++++++++++++++++++++++------ 2 files changed, 110 insertions(+), 39 deletions(-) diff --git a/odbcdol/__init__.py b/odbcdol/__init__.py index 79fe6ab..c7f3320 100644 --- a/odbcdol/__init__.py +++ b/odbcdol/__init__.py @@ -13,37 +13,63 @@ pyodbc = check_pyodbc_import() -_IDENTIFIER_RE = re.compile(r"[A-Za-z_@#][A-Za-z0-9_@#$]{0,127}") +_MAX_IDENTIFIER_LEN = 128 def _quote_identifier(name: str) -> str: - """Validate a table/column name and return it bracket-quoted for SQL Server. + """Return *name* as a bracket-quoted SQL Server identifier. - Identifiers cannot be passed as query parameters, so anything that is not a - plain (regular) SQL Server identifier is refused rather than escaped. + 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) or not _IDENTIFIER_RE.fullmatch(name): + if not isinstance(name, str): raise ValueError(f"Not a valid SQL identifier: {name!r}") - return f"[{name}]" + 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 _connection_string(**parts) -> str: - """Build an ODBC connection string, brace-quoting values that need it. +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)) - A value containing ``;``, ``{``, ``}``, ``=`` or surrounding spaces would - otherwise end its attribute and start another (e.g. a password - ``x;DATABASE=other``). Values already wrapped in braces are kept as is. + +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 _value(v) -> str: - v = str(v) - if v.startswith("{") and v.endswith("}"): - return v - if any(c in v for c in ";{}=") or v != v.strip(): - return "{" + v.replace("}", "}}") + "}" - return v - return ";".join(f"{k}={_value(v)}" for k, v in parts.items()) +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): @@ -63,7 +89,7 @@ def __init__( self.__check_dependencies() self._sql_server_client = pyodbc.connect( _connection_string( - DRIVER="{ODBC Driver 17 for SQL Server}", + DRIVER="ODBC Driver 17 for SQL Server", SERVER=f"{conn_protocol}:{host},{port}", DATABASE=db_name, UID=db_username, @@ -75,9 +101,10 @@ def __init__( self._table_name = table_name self._primary_key = primary_key - # Identifiers cannot be bound as parameters, so they are validated and - # bracket-quoted; every VALUE is bound (``?``) and never formatted in. - table = _quote_identifier(self._table_name) + # 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}});" @@ -112,7 +139,7 @@ def __check_dependencies(): ) def __getitem__(self, k): - self._cursor.execute(self._select_query, 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 @@ -125,7 +152,7 @@ def __setitem__(self, k, v): attributes=",".join(_quote_identifier(c) for c in columns), values=",".join("?" for _ in columns), ), - *(v[c] for c in columns), + tuple(_bindable(v[c]) for c in columns), ) except pyodbc.IntegrityError as e: # TODO: Raise a proper exception here @@ -134,7 +161,7 @@ def __setitem__(self, k, v): self._sql_server_client.commit() def __delitem__(self, k): - self._cursor.execute(self._del_query, 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 index 118192e..6301a1e 100644 --- a/odbcdol/tests/test_injection.py +++ b/odbcdol/tests/test_injection.py @@ -13,8 +13,8 @@ class _Cursor: def __init__(self, log): self.log = log - def execute(self, sql, *params): - self.log.append((sql, params)) + def execute(self, sql, params=()): + self.log.append((sql, tuple(params))) def fetchone(self): return ("row",) @@ -79,26 +79,70 @@ def test_values_are_bound_on_write(persister_cls, value): assert params == ("3", value) -@pytest.mark.parametrize("column", ["name) VALUES (1); --", "a b", "x]", ""]) -def test_hostile_column_names_are_refused(persister_cls, column): +@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( - "kw", [{"table_name": "person; DROP TABLE x"}, {"primary_key": "id = id OR 1"}] + "table, quoted", + [ + ("person", "[person]"), + ("dbo.person", "[dbo].[person]"), + ("[dbo].[person]", "[dbo].[person]"), + ("person]; DROP TABLE x; --", "[person]]; DROP TABLE x; --]"), + ], ) -def test_hostile_identifiers_are_refused_at_construction(persister_cls, kw): - cls, _ = persister_cls - with pytest.raises(ValueError): - cls(**kw) +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] = ?;" -def test_connection_string_values_cannot_add_attributes(persister_cls): +@pytest.mark.parametrize("key", [(1, 2), [1], {"a": 1}]) +def test_container_keys_bind_as_one_value(persister_cls, key): cls, log = persister_cls - s = cls(db_pass="pw;DATABASE=master", db_name="db") + 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 "PWD={pw;DATABASE=master}" in conn_str - assert "DATABASE=db;" in 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, "" + )