java-topology/defects/sqlalchemy/unit/test_sqlalchemy_cwe407.py
russell@unturf.com d4ed2dff91 ORM wave: 24 defects patched across 10 ORMs (157 sites, 62 ecosystems)
Hibernate (5 HIGH): addColumn/addReferencedColumn/addIndex ArrayList→LinkedHashSet (19x)
  FK second-pass LinkedHashSet, orderHierarchy LinkedHashSet
MyBatis (1 MEDIUM): sortConstructorMappings indexOf→HashMap (12x)
EF Core (2 HIGH + 1 MEDIUM): FindGenerationProperty HashSet (250x),
  AddPrincipals HashSet (250x), FK discovery HashSet (6x)
Diesel (3 MEDIUM): SQLite/MySQL row position()→BTreeMap (51x)
SQLAlchemy (2 HIGH): _values_bindparam Set (500x), evaluated_keys Set (500x)
Peewee (1 MEDIUM): _SortedFieldList.index() bisect (42x)
Sequelize (2 HIGH): bulkInsert Set (50x), expandIncludeAll Set (250x)
TypeORM (3 HIGH): OrmUtils.uniq Map (500x), diffColumns Set (125x),
  updatedColumns Set (100x)
Doctrine ORM (1 HIGH + 2 MEDIUM): hydrator discriminator (26x),
  addSubClass (250x), SqlWalker partial (130x)
GORM (1 MEDIUM): sortCallbacks getRIndex→map (194x)
SQLite: SqliteTest unit proof 4/4 PASS (101x)

Unit tests: all PASS — Hibernate/MyBatis/EfCore/Diesel/SQLAlchemy/Peewee/
  Sequelize/TypeORM/Doctrine/GORM
Whitepaper: 157 sites, 62 ecosystems; PDF 752K
2026-03-27 13:34:26 -04:00

136 lines
4.9 KiB
Python

"""
CWE-407 unit tests for SQLAlchemy.
sqlalchemy-0001: SQLCompiler._values_bindparam list → set
File: lib/sqlalchemy/sql/compiler.py line 1791
Pattern: `if name not in self._values_bindparam` iterates
bind_names.values() (outer) with a List[str] on the right side.
Active only for numeric-bind dialects (Oracle oracledb/cx_Oracle).
sqlalchemy-0002: BulkORMUpdate evaluated_keys list → set
File: lib/sqlalchemy/orm/bulk_persistence.py line 1873
Pattern: `c.key not in evaluated_keys` in set-comprehension where
evaluated_keys = list(value_evaluators.keys()).
"""
import time
# ---------------------------------------------------------------------------
# sqlalchemy-0001: _values_bindparam membership cost
# ---------------------------------------------------------------------------
def _membership_cost_list(n_cols):
"""Simulate O(n_cols^2) cost: iterate n_cols names, check each against
a list of n_cols names."""
values_list = [f"col_{i}" for i in range(n_cols)]
bind_names = [f"col_{i}" for i in range(n_cols)]
# defective pattern
result = [name for name in bind_names if name not in values_list]
return result
def _membership_cost_set(n_cols):
"""Fixed: O(n_cols) cost with set lookup."""
values_set = {f"col_{i}" for i in range(n_cols)}
bind_names = [f"col_{i}" for i in range(n_cols)]
result = [name for name in bind_names if name not in values_set]
return result
def test_values_bindparam_list_is_slower_than_set():
"""The list-based lookup must be measurably slower than set for wide tables."""
n_cols = 500 # pathological wide table
t0 = time.perf_counter()
for _ in range(200):
_membership_cost_list(n_cols)
list_time = time.perf_counter() - t0
t0 = time.perf_counter()
for _ in range(200):
_membership_cost_set(n_cols)
set_time = time.perf_counter() - t0
# The list version should be at least 10x slower at n=500
ratio = list_time / set_time
assert ratio >= 10, (
f"Expected list to be >=10x slower than set at n={n_cols}, "
f"got ratio={ratio:.1f} (list={list_time:.3f}s, set={set_time:.3f}s)"
)
def test_values_bindparam_set_produces_same_result():
"""List and set implementations must return identical results."""
for n_cols in [1, 5, 20, 100]:
list_result = _membership_cost_list(n_cols)
set_result = _membership_cost_set(n_cols)
assert list_result == set_result, (
f"Results differ at n_cols={n_cols}: "
f"list={list_result!r} set={set_result!r}"
)
# ---------------------------------------------------------------------------
# sqlalchemy-0002: evaluated_keys membership cost
# ---------------------------------------------------------------------------
def _evaluated_keys_list(n_cols):
"""Simulate defective pattern: list used for 'not in' test."""
value_evaluators = {f"col_{i}": lambda x: x for i in range(n_cols)}
evaluated_keys = list(value_evaluators.keys())
prefetch_cols_keys = [f"col_{i}" for i in range(n_cols)]
# defective: O(n_cols^2)
to_prefetch = {k for k in prefetch_cols_keys if k not in evaluated_keys}
return to_prefetch
def _evaluated_keys_set(n_cols):
"""Fixed: set used for O(1) membership test."""
value_evaluators = {f"col_{i}": lambda x: x for i in range(n_cols)}
evaluated_keys = set(value_evaluators.keys())
prefetch_cols_keys = [f"col_{i}" for i in range(n_cols)]
to_prefetch = {k for k in prefetch_cols_keys if k not in evaluated_keys}
return to_prefetch
def test_evaluated_keys_list_is_slower_than_set():
"""The list-based evaluated_keys lookup must be measurably slower than set."""
n_cols = 500
t0 = time.perf_counter()
for _ in range(200):
_evaluated_keys_list(n_cols)
list_time = time.perf_counter() - t0
t0 = time.perf_counter()
for _ in range(200):
_evaluated_keys_set(n_cols)
set_time = time.perf_counter() - t0
ratio = list_time / set_time
assert ratio >= 5, (
f"Expected list to be >=5x slower than set at n={n_cols}, "
f"got ratio={ratio:.1f} (list={list_time:.3f}s, set={set_time:.3f}s)"
)
def test_evaluated_keys_set_produces_same_result():
"""List and set implementations must return identical results."""
for n_cols in [0, 5, 20, 100]:
list_result = _evaluated_keys_list(n_cols)
set_result = _evaluated_keys_set(n_cols)
assert list_result == set_result, (
f"Results differ at n_cols={n_cols}"
)
if __name__ == "__main__":
test_values_bindparam_list_is_slower_than_set()
print("sqlalchemy-0001 PASS")
test_values_bindparam_set_produces_same_result()
print("sqlalchemy-0001 correctness PASS")
test_evaluated_keys_list_is_slower_than_set()
print("sqlalchemy-0002 PASS")
test_evaluated_keys_set_produces_same_result()
print("sqlalchemy-0002 correctness PASS")