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
110 changes: 83 additions & 27 deletions odbcdol/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
"""

import platform
import re
import subprocess
from collections.abc import MutableMapping

Expand All @@ -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,
Expand All @@ -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():
Expand Down Expand Up @@ -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
Expand All @@ -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):
Expand Down
148 changes: 148 additions & 0 deletions odbcdol/tests/test_injection.py
Original file line number Diff line number Diff line change
@@ -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, ""
)
Loading