Team Ai
Datasetpublic

codekingpro/portable-devtools

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