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