codekingpro/portable-devtools
114k
1# testing/fixtures/sql.py
2# Copyright (C) 2005-2024 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
13import sys
14
15import sqlalchemy as sa
16from .base import TestBase
17from .. import config
18from .. import mock
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 savepoints = getattr(config.requirements, "savepoints", False)
137 if savepoints:
138 savepoints = savepoints.enabled
139
140 # no need to run deletes if tables are recreated on setup
141 if (
142 self.run_define_tables != "each"
143 and self.run_create_tables != "each"
144 and self.run_deletes == "each"
145 ):
146 with self.bind.begin() as conn:
147 for table in reversed(
148 [
149 t
150 for (t, fks) in sort_tables_and_constraints(
151 self._tables_metadata.tables.values()
152 )
153 if t is not None
154 ]
155 ):
156 try:
157 if savepoints:
158 with conn.begin_nested():
159 conn.execute(table.delete())
160 else:
161 conn.execute(table.delete())
162 except sa.exc.DBAPIError as ex:
163 print(
164 ("Error emptying table %s: %r" % (table, ex)),
165 file=sys.stderr,
166 )
167
168 @classmethod
169 def _teardown_once_metadata_bind(cls):
170 if cls.run_create_tables:
171 drop_all_tables_from_metadata(cls._tables_metadata, cls.bind)
172
173 if cls.run_dispose_bind == "once":
174 cls.dispose_bind(cls.bind)
175
176 cls._tables_metadata.bind = None
177
178 if cls.run_setup_bind is not None:
179 cls.bind = None
180
181 @classmethod
182 def setup_bind(cls):
183 return config.db
184
185 @classmethod
186 def dispose_bind(cls, bind):
187 if hasattr(bind, "dispose"):
188 bind.dispose()
189 elif hasattr(bind, "close"):
190 bind.close()
191
192 @classmethod
193 def define_tables(cls, metadata):
194 pass
195
196 @classmethod
197 def fixtures(cls):
198 return {}
199
200 @classmethod
201 def insert_data(cls, connection):
202 pass
203
204 def sql_count_(self, count, fn):
205 self.assert_sql_count(self.bind, fn, count)
206
207 def sql_eq_(self, callable_, statements):
208 self.assert_sql(self.bind, callable_, statements)
209
210 @classmethod
211 def _load_fixtures(cls):
212 """Insert rows as represented by the fixtures() method."""
213 headers, rows = {}, {}
214 for table, data in cls.fixtures().items():
215 if len(data) < 2:
216 continue
217 if isinstance(table, str):
218 table = cls.tables[table]
219 headers[table] = data[0]
220 rows[table] = data[1:]
221 for table, fks in sort_tables_and_constraints(
222 cls._tables_metadata.tables.values()
223 ):
224 if table is None:
225 continue
226 if table not in headers:
227 continue
228 with cls.bind.begin() as conn:
229 conn.execute(
230 table.insert(),
231 [
232 dict(zip(headers[table], column_values))
233 for column_values in rows[table]
234 ],
235 )
236
237
238class NoCache:
239 @config.fixture(autouse=True, scope="function")
240 def _disable_cache(self):
241 _cache = config.db._compiled_cache
242 config.db._compiled_cache = None
243 yield
244 config.db._compiled_cache = _cache
245
246
247class RemovesEvents:
248 @util.memoized_property
249 def _event_fns(self):
250 return set()
251
252 def event_listen(self, target, name, fn, **kw):
253 self._event_fns.add((target, name, fn))
254 event.listen(target, name, fn, **kw)
255
256 @config.fixture(autouse=True, scope="function")
257 def _remove_events(self):
258 yield
259 for key in self._event_fns:
260 event.remove(*key)
261
262
263class ComputedReflectionFixtureTest(TablesTest):
264 run_inserts = run_deletes = None
265
266 __backend__ = True
267 __requires__ = ("computed_columns", "table_reflection")
268
269 regexp = re.compile(r"[\[\]\(\)\s`'\"]*")
270
271 def normalize(self, text):
272 return self.regexp.sub("", text).lower()
273
274 @classmethod
275 def define_tables(cls, metadata):
276 from ... import Integer
277 from ... import testing
278 from ...schema import Column
279 from ...schema import Computed
280 from ...schema import Table
281
282 Table(
283 "computed_default_table",
284 metadata,
285 Column("id", Integer, primary_key=True),
286 Column("normal", Integer),
287 Column("computed_col", Integer, Computed("normal + 42")),
288 Column("with_default", Integer, server_default="42"),
289 )
290
291 t = Table(
292 "computed_column_table",
293 metadata,
294 Column("id", Integer, primary_key=True),
295 Column("normal", Integer),
296 Column("computed_no_flag", Integer, Computed("normal + 42")),
297 )
298
299 if testing.requires.schemas.enabled:
300 t2 = Table(
301 "computed_column_table",
302 metadata,
303 Column("id", Integer, primary_key=True),
304 Column("normal", Integer),
305 Column("computed_no_flag", Integer, Computed("normal / 42")),
306 schema=config.test_schema,
307 )
308
309 if testing.requires.computed_columns_virtual.enabled:
310 t.append_column(
311 Column(
312 "computed_virtual",
313 Integer,
314 Computed("normal + 2", persisted=False),
315 )
316 )
317 if testing.requires.schemas.enabled:
318 t2.append_column(
319 Column(
320 "computed_virtual",
321 Integer,
322 Computed("normal / 2", persisted=False),
323 )
324 )
325 if testing.requires.computed_columns_stored.enabled:
326 t.append_column(
327 Column(
328 "computed_stored",
329 Integer,
330 Computed("normal - 42", persisted=True),
331 )
332 )
333 if testing.requires.schemas.enabled:
334 t2.append_column(
335 Column(
336 "computed_stored",
337 Integer,
338 Computed("normal * 42", persisted=True),
339 )
340 )
341
342
343class CacheKeyFixture:
344 def _compare_equal(self, a, b, compare_values):
345 a_key = a._generate_cache_key()
346 b_key = b._generate_cache_key()
347
348 if a_key is None:
349 assert a._annotations.get("nocache")
350
351 assert b_key is None
352 else:
353 eq_(a_key.key, b_key.key)
354 eq_(hash(a_key.key), hash(b_key.key))
355
356 for a_param, b_param in zip(a_key.bindparams, b_key.bindparams):
357 assert a_param.compare(b_param, compare_values=compare_values)
358 return a_key, b_key
359
360 def _run_cache_key_fixture(self, fixture, compare_values):
361 case_a = fixture()
362 case_b = fixture()
363
364 for a, b in itertools.combinations_with_replacement(
365 range(len(case_a)), 2
366 ):
367 if a == b:
368 a_key, b_key = self._compare_equal(
369 case_a[a], case_b[b], compare_values
370 )
371 if a_key is None:
372 continue
373 else:
374 a_key = case_a[a]._generate_cache_key()
375 b_key = case_b[b]._generate_cache_key()
376
377 if a_key is None or b_key is None:
378 if a_key is None:
379 assert case_a[a]._annotations.get("nocache")
380 if b_key is None:
381 assert case_b[b]._annotations.get("nocache")
382 continue
383
384 if a_key.key == b_key.key:
385 for a_param, b_param in zip(
386 a_key.bindparams, b_key.bindparams
387 ):
388 if not a_param.compare(
389 b_param, compare_values=compare_values
390 ):
391 break
392 else:
393 # this fails unconditionally since we could not
394 # find bound parameter values that differed.
395 # Usually we intended to get two distinct keys here
396 # so the failure will be more descriptive using the
397 # ne_() assertion.
398 ne_(a_key.key, b_key.key)
399 else:
400 ne_(a_key.key, b_key.key)
401
402 # ClauseElement-specific test to ensure the cache key
403 # collected all the bound parameters that aren't marked
404 # as "literal execute"
405 if isinstance(case_a[a], ClauseElement) and isinstance(
406 case_b[b], ClauseElement
407 ):
408 assert_a_params = []
409 assert_b_params = []
410
411 for elem in visitors.iterate(case_a[a]):
412 if elem.__visit_name__ == "bindparam":
413 assert_a_params.append(elem)
414
415 for elem in visitors.iterate(case_b[b]):
416 if elem.__visit_name__ == "bindparam":
417 assert_b_params.append(elem)
418
419 # note we're asserting the order of the params as well as
420 # if there are dupes or not. ordering has to be
421 # deterministic and matches what a traversal would provide.
422 eq_(
423 sorted(a_key.bindparams, key=lambda b: b.key),
424 sorted(
425 util.unique_list(assert_a_params), key=lambda b: b.key
426 ),
427 )
428 eq_(
429 sorted(b_key.bindparams, key=lambda b: b.key),
430 sorted(
431 util.unique_list(assert_b_params), key=lambda b: b.key
432 ),
433 )
434
435 def _run_cache_key_equal_fixture(self, fixture, compare_values):
436 case_a = fixture()
437 case_b = fixture()
438
439 for a, b in itertools.combinations_with_replacement(
440 range(len(case_a)), 2
441 ):
442 self._compare_equal(case_a[a], case_b[b], compare_values)
443
444
445def insertmanyvalues_fixture(
446 connection, randomize_rows=False, warn_on_downgraded=False
447):
448 dialect = connection.dialect
449 orig_dialect = dialect._deliver_insertmanyvalues_batches
450 orig_conn = connection._exec_insertmany_context
451
452 class RandomCursor:
453 __slots__ = ("cursor",)
454
455 def __init__(self, cursor):
456 self.cursor = cursor
457
458 # only this method is called by the deliver method.
459 # by not having the other methods we assert that those aren't being
460 # used
461
462 @property
463 def description(self):
464 return self.cursor.description
465
466 def fetchall(self):
467 rows = self.cursor.fetchall()
468 rows = list(rows)
469 random.shuffle(rows)
470 return rows
471
472 def _deliver_insertmanyvalues_batches(
473 cursor, statement, parameters, generic_setinputsizes, context
474 ):
475 if randomize_rows:
476 cursor = RandomCursor(cursor)
477 for batch in orig_dialect(
478 cursor, statement, parameters, generic_setinputsizes, context
479 ):
480 if warn_on_downgraded and batch.is_downgraded:
481 util.warn("Batches were downgraded for sorted INSERT")
482
483 yield batch
484
485 def _exec_insertmany_context(dialect, context):
486 with mock.patch.object(
487 dialect,
488 "_deliver_insertmanyvalues_batches",
489 new=_deliver_insertmanyvalues_batches,
490 ):
491 return orig_conn(dialect, context)
492
493 connection._exec_insertmany_context = _exec_insertmany_context
494 