codekingpro/portable-devtools
114k
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 