codekingpro/portable-devtools
115k
1# testing/fixtures/sql.py
2# Copyright (C) 2005-2026 the SQLAlchemy authors and contributors
3# <see AUTHORS file>
4#
5# This module is part of SQLAlchemy and is released under
6# the MIT License: https://www.opensource.org/licenses/mit-license.php
7# mypy: ignore-errors
8from __future__ import annotations
9
10import itertools
11import random
12import re
13
14import sqlalchemy as sa
15from .base import TestBase
16from .. import config
17from .. import mock
18from .. import provision
19from ..assertions import eq_
20from ..assertions import ne_
21from ..util import adict
22from ..util import drop_all_tables_from_metadata
23from ... import event
24from ... import util
25from ...schema import sort_tables_and_constraints
26from ...sql import visitors
27from ...sql.elements import ClauseElement
28
29
30class TablesTest(TestBase):
31 # 'once', None
32 run_setup_bind = "once"
33
34 # 'once', 'each', None
35 run_define_tables = "once"
36
37 # 'once', 'each', None
38 run_create_tables = "once"
39
40 # 'once', 'each', None
41 run_inserts = "each"
42
43 # 'each', None
44 run_deletes = "each"
45
46 # 'once', None
47 run_dispose_bind = None
48
49 bind = None
50 _tables_metadata = None
51 tables = None
52 other = None
53 sequences = None
54
55 @config.fixture(autouse=True, scope="class")
56 def _setup_tables_test_class(self):
57 cls = self.__class__
58 cls._init_class()
59
60 cls._setup_once_tables()
61
62 cls._setup_once_inserts()
63
64 yield
65
66 cls._teardown_once_metadata_bind()
67
68 @config.fixture(autouse=True, scope="function")
69 def _setup_tables_test_instance(self):
70 self._setup_each_tables()
71 self._setup_each_inserts()
72
73 yield
74
75 self._teardown_each_tables()
76
77 @property
78 def tables_test_metadata(self):
79 return self._tables_metadata
80
81 @classmethod
82 def _init_class(cls):
83 if cls.run_define_tables == "each":
84 if cls.run_create_tables == "once":
85 cls.run_create_tables = "each"
86 assert cls.run_inserts in ("each", None)
87
88 cls.other = adict()
89 cls.tables = adict()
90 cls.sequences = adict()
91
92 cls.bind = cls.setup_bind()
93 cls._tables_metadata = sa.MetaData()
94
95 @classmethod
96 def _setup_once_inserts(cls):
97 if cls.run_inserts == "once":
98 cls._load_fixtures()
99 with cls.bind.begin() as conn:
100 cls.insert_data(conn)
101
102 @classmethod
103 def _setup_once_tables(cls):
104 if cls.run_define_tables == "once":
105 cls.define_tables(cls._tables_metadata)
106 if cls.run_create_tables == "once":
107 cls._tables_metadata.create_all(cls.bind)
108 cls.tables.update(cls._tables_metadata.tables)
109 cls.sequences.update(cls._tables_metadata._sequences)
110
111 def _setup_each_tables(self):
112 if self.run_define_tables == "each":
113 self.define_tables(self._tables_metadata)
114 if self.run_create_tables == "each":
115 self._tables_metadata.create_all(self.bind)
116 self.tables.update(self._tables_metadata.tables)
117 self.sequences.update(self._tables_metadata._sequences)
118 elif self.run_create_tables == "each":
119 self._tables_metadata.create_all(self.bind)
120
121 def _setup_each_inserts(self):
122 if self.run_inserts == "each":
123 self._load_fixtures()
124 with self.bind.begin() as conn:
125 self.insert_data(conn)
126
127 def _teardown_each_tables(self):
128 if self.run_define_tables == "each":
129 self.tables.clear()
130 if self.run_create_tables == "each":
131 drop_all_tables_from_metadata(self._tables_metadata, self.bind)
132 self._tables_metadata.clear()
133 elif self.run_create_tables == "each":
134 drop_all_tables_from_metadata(self._tables_metadata, self.bind)
135
136 # no need to run deletes if tables are recreated on setup
137 if (
138 self.run_define_tables != "each"
139 and self.run_create_tables == "once"
140 and self.run_deletes == "each"
141 ):
142 with self.bind.begin() as conn:
143 provision.delete_from_all_tables(
144 conn, config, self._tables_metadata
145 )
146
147 @classmethod
148 def _teardown_once_metadata_bind(cls):
149 if cls.run_create_tables:
150 drop_all_tables_from_metadata(cls._tables_metadata, cls.bind)
151
152 if cls.run_dispose_bind == "once":
153 cls.dispose_bind(cls.bind)
154
155 cls._tables_metadata.bind = None
156
157 if cls.run_setup_bind is not None:
158 cls.bind = None
159
160 @classmethod
161 def setup_bind(cls):
162 return config.db
163
164 @classmethod
165 def dispose_bind(cls, bind):
166 if hasattr(bind, "dispose"):
167 bind.dispose()
168 elif hasattr(bind, "close"):
169 bind.close()
170
171 @classmethod
172 def define_tables(cls, metadata):
173 pass
174
175 @classmethod
176 def fixtures(cls):
177 return {}
178
179 @classmethod
180 def insert_data(cls, connection):
181 pass
182
183 def sql_count_(self, count, fn):
184 self.assert_sql_count(self.bind, fn, count)
185
186 def sql_eq_(self, callable_, statements):
187 self.assert_sql(self.bind, callable_, statements)
188
189 @classmethod
190 def _load_fixtures(cls):
191 """Insert rows as represented by the fixtures() method."""
192 headers, rows = {}, {}
193 for table, data in cls.fixtures().items():
194 if len(data) < 2:
195 continue
196 if isinstance(table, str):
197 table = cls.tables[table]
198 headers[table] = data[0]
199 rows[table] = data[1:]
200 for table, fks in sort_tables_and_constraints(
201 cls._tables_metadata.tables.values()
202 ):
203 if table is None:
204 continue
205 if table not in headers:
206 continue
207 with cls.bind.begin() as conn:
208 conn.execute(
209 table.insert(),
210 [
211 dict(zip(headers[table], column_values))
212 for column_values in rows[table]
213 ],
214 )
215
216
217class NoCache:
218 @config.fixture(autouse=True, scope="function")
219 def _disable_cache(self):
220 _cache = config.db._compiled_cache
221 config.db._compiled_cache = None
222 yield
223 config.db._compiled_cache = _cache
224
225
226class RemovesEvents:
227 @util.memoized_property
228 def _event_fns(self):
229 return set()
230
231 def event_listen(self, target, name, fn, **kw):
232 self._event_fns.add((target, name, fn))
233 event.listen(target, name, fn, **kw)
234
235 @config.fixture(autouse=True, scope="function")
236 def _remove_events(self):
237 yield
238 for key in self._event_fns:
239 event.remove(*key)
240
241
242class ComputedReflectionFixtureTest(TablesTest):
243 run_inserts = run_deletes = None
244
245 __backend__ = True
246 __requires__ = ("computed_columns", "table_reflection")
247
248 regexp = re.compile(r"[\[\]\(\)\s`'\"]*")
249
250 def normalize(self, text):
251 return self.regexp.sub("", text).lower()
252
253 @classmethod
254 def define_tables(cls, metadata):
255 from ... import Integer
256 from ... import testing
257 from ...schema import Column
258 from ...schema import Computed
259 from ...schema import Table
260
261 Table(
262 "computed_default_table",
263 metadata,
264 Column("id", Integer, primary_key=True),
265 Column("normal", Integer),
266 Column("computed_col", Integer, Computed("normal + 42")),
267 Column("with_default", Integer, server_default="42"),
268 )
269
270 t = Table(
271 "computed_column_table",
272 metadata,
273 Column("id", Integer, primary_key=True),
274 Column("normal", Integer),
275 Column("computed_no_flag", Integer, Computed("normal + 42")),
276 )
277
278 if testing.requires.schemas.enabled:
279 t2 = Table(
280 "computed_column_table",
281 metadata,
282 Column("id", Integer, primary_key=True),
283 Column("normal", Integer),
284 Column("computed_no_flag", Integer, Computed("normal / 42")),
285 schema=config.test_schema,
286 )
287
288 if testing.requires.computed_columns_virtual.enabled:
289 t.append_column(
290 Column(
291 "computed_virtual",
292 Integer,
293 Computed("normal + 2", persisted=False),
294 )
295 )
296 if testing.requires.schemas.enabled:
297 t2.append_column(
298 Column(
299 "computed_virtual",
300 Integer,
301 Computed("normal / 2", persisted=False),
302 )
303 )
304 if testing.requires.computed_columns_stored.enabled:
305 t.append_column(
306 Column(
307 "computed_stored",
308 Integer,
309 Computed("normal - 42", persisted=True),
310 )
311 )
312 if testing.requires.schemas.enabled:
313 t2.append_column(
314 Column(
315 "computed_stored",
316 Integer,
317 Computed("normal * 42", persisted=True),
318 )
319 )
320
321
322class CacheKeyFixture:
323 def _compare_equal(self, a, b, compare_values):
324 a_key = a._generate_cache_key()
325 b_key = b._generate_cache_key()
326
327 if a_key is None:
328 assert a._annotations.get("nocache")
329
330 assert b_key is None
331 else:
332 eq_(a_key.key, b_key.key)
333 eq_(hash(a_key.key), hash(b_key.key))
334
335 for a_param, b_param in zip(a_key.bindparams, b_key.bindparams):
336 assert a_param.compare(b_param, compare_values=compare_values)
337 return a_key, b_key
338
339 def _run_cache_key_fixture(self, fixture, compare_values):
340 case_a = fixture()
341 case_b = fixture()
342
343 for a, b in itertools.combinations_with_replacement(
344 range(len(case_a)), 2
345 ):
346 if a == b:
347 a_key, b_key = self._compare_equal(
348 case_a[a], case_b[b], compare_values
349 )
350 if a_key is None:
351 continue
352 else:
353 a_key = case_a[a]._generate_cache_key()
354 b_key = case_b[b]._generate_cache_key()
355
356 if a_key is None or b_key is None:
357 if a_key is None:
358 assert case_a[a]._annotations.get("nocache")
359 if b_key is None:
360 assert case_b[b]._annotations.get("nocache")
361 continue
362
363 if a_key.key == b_key.key:
364 for a_param, b_param in zip(
365 a_key.bindparams, b_key.bindparams
366 ):
367 if not a_param.compare(
368 b_param, compare_values=compare_values
369 ):
370 break
371 else:
372 # this fails unconditionally since we could not
373 # find bound parameter values that differed.
374 # Usually we intended to get two distinct keys here
375 # so the failure will be more descriptive using the
376 # ne_() assertion.
377 ne_(a_key.key, b_key.key)
378 else:
379 ne_(a_key.key, b_key.key)
380
381 # ClauseElement-specific test to ensure the cache key
382 # collected all the bound parameters that aren't marked
383 # as "literal execute"
384 if isinstance(case_a[a], ClauseElement) and isinstance(
385 case_b[b], ClauseElement
386 ):
387 assert_a_params = []
388 assert_b_params = []
389
390 for elem in visitors.iterate(case_a[a]):
391 if elem.__visit_name__ == "bindparam":
392 assert_a_params.append(elem)
393
394 for elem in visitors.iterate(case_b[b]):
395 if elem.__visit_name__ == "bindparam":
396 assert_b_params.append(elem)
397
398 # note we're asserting the order of the params as well as
399 # if there are dupes or not. ordering has to be
400 # deterministic and matches what a traversal would provide.
401 eq_(
402 sorted(a_key.bindparams, key=lambda b: b.key),
403 sorted(
404 util.unique_list(assert_a_params), key=lambda b: b.key
405 ),
406 )
407 eq_(
408 sorted(b_key.bindparams, key=lambda b: b.key),
409 sorted(
410 util.unique_list(assert_b_params), key=lambda b: b.key
411 ),
412 )
413
414 def _run_cache_key_equal_fixture(self, fixture, compare_values):
415 case_a = fixture()
416 case_b = fixture()
417
418 for a, b in itertools.combinations_with_replacement(
419 range(len(case_a)), 2
420 ):
421 self._compare_equal(case_a[a], case_b[b], compare_values)
422
423
424def insertmanyvalues_fixture(
425 connection, randomize_rows=False, warn_on_downgraded=False
426):
427 dialect = connection.dialect
428 orig_dialect = dialect._deliver_insertmanyvalues_batches
429 orig_conn = connection._exec_insertmany_context
430
431 class RandomCursor:
432 __slots__ = ("cursor",)
433
434 def __init__(self, cursor):
435 self.cursor = cursor
436
437 # only this method is called by the deliver method.
438 # by not having the other methods we assert that those aren't being
439 # used
440
441 @property
442 def description(self):
443 return self.cursor.description
444
445 def fetchall(self):
446 rows = self.cursor.fetchall()
447 rows = list(rows)
448 random.shuffle(rows)
449 return rows
450
451 def _deliver_insertmanyvalues_batches(
452 connection,
453 cursor,
454 statement,
455 parameters,
456 generic_setinputsizes,
457 context,
458 ):
459 if randomize_rows:
460 cursor = RandomCursor(cursor)
461 for batch in orig_dialect(
462 connection,
463 cursor,
464 statement,
465 parameters,
466 generic_setinputsizes,
467 context,
468 ):
469 if warn_on_downgraded and batch.is_downgraded:
470 util.warn("Batches were downgraded for sorted INSERT")
471
472 yield batch
473
474 def _exec_insertmany_context(dialect, context):
475 with mock.patch.object(
476 dialect,
477 "_deliver_insertmanyvalues_batches",
478 new=_deliver_insertmanyvalues_batches,
479 ):
480 return orig_conn(dialect, context)
481
482 connection._exec_insertmany_context = _exec_insertmany_context
483 