Team Ai
Datasetpublic

codekingpro/portable-devtools

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