""" 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")