Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
sql.py494 linesDownload Raw Back to fixtures
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 
codekingpro/portable-devtools · Team Ai