Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
base.py385 linesDownload Raw Back to fixtures
1# testing/fixtures/base.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
8
9
10from __future__ import annotations
11
12import sqlalchemy as sa
13from .. import assertions
14from .. import config
15from ..assertions import eq_
16from ..util import drop_all_tables_from_metadata
17from ..util import picklers
18from ... import Column
19from ... import func
20from ... import Integer
21from ... import select
22from ... import Table
23from ...orm import DeclarativeBase
24from ...orm import MappedAsDataclass
25from ...orm import registry
26
27
28@config.mark_base_test_class()
29class TestBase:
30    # A sequence of requirement names matching testing.requires decorators
31    __requires__ = ()
32
33    # A sequence of dialect names to exclude from the test class.
34    __unsupported_on__ = ()
35
36    # If present, test class is only runnable for the *single* specified
37    # dialect.  If you need multiple, use __unsupported_on__ and invert.
38    __only_on__ = None
39
40    # A sequence of no-arg callables. If any are True, the entire testcase is
41    # skipped.
42    __skip_if__ = None
43
44    # if True, the testing reaper will not attempt to touch connection
45    # state after a test is completed and before the outer teardown
46    # starts
47    __leave_connections_for_teardown__ = False
48
49    def assert_(self, val, msg=None):
50        assert val, msg
51
52    @config.fixture()
53    def nocache(self):
54        _cache = config.db._compiled_cache
55        config.db._compiled_cache = None
56        yield
57        config.db._compiled_cache = _cache
58
59    @config.fixture()
60    def connection_no_trans(self):
61        eng = getattr(self, "bind", None) or config.db
62
63        with eng.connect() as conn:
64            yield conn
65
66    @config.fixture()
67    def connection(self):
68        global _connection_fixture_connection
69
70        eng = getattr(self, "bind", None) or config.db
71
72        conn = eng.connect()
73        trans = conn.begin()
74
75        _connection_fixture_connection = conn
76        yield conn
77
78        _connection_fixture_connection = None
79
80        if trans.is_active:
81            trans.rollback()
82        # trans would not be active here if the test is using
83        # the legacy @provide_metadata decorator still, as it will
84        # run a close all connections.
85        conn.close()
86
87    @config.fixture()
88    def close_result_when_finished(self):
89        to_close = []
90        to_consume = []
91
92        def go(result, consume=False):
93            to_close.append(result)
94            if consume:
95                to_consume.append(result)
96
97        yield go
98        for r in to_consume:
99            try:
100                r.all()
101            except:
102                pass
103        for r in to_close:
104            try:
105                r.close()
106            except:
107                pass
108
109    @config.fixture()
110    def registry(self, metadata):
111        reg = registry(
112            metadata=metadata,
113            type_annotation_map={
114                str: sa.String().with_variant(
115                    sa.String(50), "mysql", "mariadb", "oracle"
116                )
117            },
118        )
119        yield reg
120        reg.dispose()
121
122    @config.fixture
123    def decl_base(self, metadata):
124        _md = metadata
125
126        class Base(DeclarativeBase):
127            metadata = _md
128            type_annotation_map = {
129                str: sa.String().with_variant(
130                    sa.String(50), "mysql", "mariadb", "oracle"
131                )
132            }
133
134        yield Base
135        Base.registry.dispose()
136
137    @config.fixture
138    def dc_decl_base(self, metadata):
139        _md = metadata
140
141        class Base(MappedAsDataclass, DeclarativeBase):
142            metadata = _md
143            type_annotation_map = {
144                str: sa.String().with_variant(
145                    sa.String(50), "mysql", "mariadb"
146                )
147            }
148
149        yield Base
150        Base.registry.dispose()
151
152    @config.fixture()
153    def future_connection(self, future_engine, connection):
154        # integrate the future_engine and connection fixtures so
155        # that users of the "connection" fixture will get at the
156        # "future" connection
157        yield connection
158
159    @config.fixture()
160    def future_engine(self):
161        yield
162
163    @config.fixture()
164    def testing_engine(self):
165        from .. import engines
166
167        def gen_testing_engine(
168            url=None,
169            options=None,
170            asyncio=False,
171        ):
172            if options is None:
173                options = {}
174            options["scope"] = "fixture"
175            return engines.testing_engine(
176                url=url,
177                options=options,
178                asyncio=asyncio,
179            )
180
181        yield gen_testing_engine
182
183        engines.testing_reaper._drop_testing_engines("fixture")
184
185    @config.fixture()
186    def async_testing_engine(self, testing_engine):
187        def go(**kw):
188            kw["asyncio"] = True
189            return testing_engine(**kw)
190
191        return go
192
193    @config.fixture(params=picklers())
194    def picklers(self, request):
195        yield request.param
196
197    @config.fixture()
198    def metadata(self, request):
199        """Provide bound MetaData for a single test, dropping afterwards."""
200
201        from ...sql import schema
202
203        metadata = schema.MetaData()
204        request.instance.metadata = metadata
205        yield metadata
206        del request.instance.metadata
207
208        if (
209            _connection_fixture_connection
210            and _connection_fixture_connection.in_transaction()
211        ):
212            trans = _connection_fixture_connection.get_transaction()
213            trans.rollback()
214            with _connection_fixture_connection.begin():
215                drop_all_tables_from_metadata(
216                    metadata, _connection_fixture_connection
217                )
218        else:
219            drop_all_tables_from_metadata(metadata, config.db)
220
221    @config.fixture()
222    def thirdparty_dialect(self):
223        from ...dialects import registry
224
225        name = None
226
227        def go(dialect_cls):
228            nonlocal name
229            name = dialect_cls.name
230            assert name, "name is required"
231            registry.impls[name] = dialect_cls
232            return dialect_cls
233
234        yield go
235
236        assert name is not None
237        del registry.impls[name]
238
239    @config.fixture(
240        params=[
241            (rollback, second_operation, begin_nested)
242            for rollback in (True, False)
243            for second_operation in ("none", "execute", "begin")
244            for begin_nested in (
245                True,
246                False,
247            )
248        ]
249    )
250    def trans_ctx_manager_fixture(self, request, metadata):
251        rollback, second_operation, begin_nested = request.param
252
253        t = Table("test", metadata, Column("data", Integer))
254        eng = getattr(self, "bind", None) or config.db
255
256        t.create(eng)
257
258        def run_test(subject, trans_on_subject, execute_on_subject):
259            with subject.begin() as trans:
260                if begin_nested:
261                    if not config.requirements.savepoints.enabled:
262                        config.skip_test("savepoints not enabled")
263                    if execute_on_subject:
264                        nested_trans = subject.begin_nested()
265                    else:
266                        nested_trans = trans.begin_nested()
267
268                    with nested_trans:
269                        if execute_on_subject:
270                            subject.execute(t.insert(), {"data": 10})
271                        else:
272                            trans.execute(t.insert(), {"data": 10})
273
274                        # for nested trans, we always commit/rollback on the
275                        # "nested trans" object itself.
276                        # only Session(future=False) will affect savepoint
277                        # transaction for session.commit/rollback
278
279                        if rollback:
280                            nested_trans.rollback()
281                        else:
282                            nested_trans.commit()
283
284                        if second_operation != "none":
285                            with assertions.expect_raises_message(
286                                sa.exc.InvalidRequestError,
287                                "Can't operate on closed transaction "
288                                "inside context "
289                                "manager.  Please complete the context "
290                                "manager "
291                                "before emitting further commands.",
292                            ):
293                                if second_operation == "execute":
294                                    if execute_on_subject:
295                                        subject.execute(
296                                            t.insert(), {"data": 12}
297                                        )
298                                    else:
299                                        trans.execute(t.insert(), {"data": 12})
300                                elif second_operation == "begin":
301                                    if execute_on_subject:
302                                        subject.begin_nested()
303                                    else:
304                                        trans.begin_nested()
305
306                    # outside the nested trans block, but still inside the
307                    # transaction block, we can run SQL, and it will be
308                    # committed
309                    if execute_on_subject:
310                        subject.execute(t.insert(), {"data": 14})
311                    else:
312                        trans.execute(t.insert(), {"data": 14})
313
314                else:
315                    if execute_on_subject:
316                        subject.execute(t.insert(), {"data": 10})
317                    else:
318                        trans.execute(t.insert(), {"data": 10})
319
320                    if trans_on_subject:
321                        if rollback:
322                            subject.rollback()
323                        else:
324                            subject.commit()
325                    else:
326                        if rollback:
327                            trans.rollback()
328                        else:
329                            trans.commit()
330
331                    if second_operation != "none":
332                        with assertions.expect_raises_message(
333                            sa.exc.InvalidRequestError,
334                            "Can't operate on closed transaction inside "
335                            "context "
336                            "manager.  Please complete the context manager "
337                            "before emitting further commands.",
338                        ):
339                            if second_operation == "execute":
340                                if execute_on_subject:
341                                    subject.execute(t.insert(), {"data": 12})
342                                else:
343                                    trans.execute(t.insert(), {"data": 12})
344                            elif second_operation == "begin":
345                                if hasattr(trans, "begin"):
346                                    trans.begin()
347                                else:
348                                    subject.begin()
349                            elif second_operation == "begin_nested":
350                                if execute_on_subject:
351                                    subject.begin_nested()
352                                else:
353                                    trans.begin_nested()
354
355            expected_committed = 0
356            if begin_nested:
357                # begin_nested variant, we inserted a row after the nested
358                # block
359                expected_committed += 1
360            if not rollback:
361                # not rollback variant, our row inserted in the target
362                # block itself would be committed
363                expected_committed += 1
364
365            if execute_on_subject:
366                eq_(
367                    subject.scalar(select(func.count()).select_from(t)),
368                    expected_committed,
369                )
370            else:
371                with subject.connect() as conn:
372                    eq_(
373                        conn.scalar(select(func.count()).select_from(t)),
374                        expected_committed,
375                    )
376
377        return run_test
378
379
380_connection_fixture_connection = None
381
382
383class FutureEngineMixin:
384    """alembic's suite still using this"""
385 
codekingpro/portable-devtools · Team Ai