codekingpro/portable-devtools
115k
1# testing/engines.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 collections
13import re
14import typing
15from typing import Any
16from typing import Dict
17from typing import Optional
18from typing import Union
19import warnings
20import weakref
21
22from . import config
23from .util import decorator
24from .util import gc_collect
25from .. import event
26from .. import pool
27from ..util import await_only
28from ..util.typing import Literal
29
30
31if typing.TYPE_CHECKING:
32 from ..engine import Engine
33 from ..engine.url import URL
34 from ..ext.asyncio import AsyncEngine
35
36
37class ConnectionKiller:
38 def __init__(self):
39 self.proxy_refs = weakref.WeakKeyDictionary()
40 self.testing_engines = collections.defaultdict(set)
41 self.dbapi_connections = set()
42
43 def add_pool(self, pool):
44 event.listen(pool, "checkout", self._add_conn)
45 event.listen(pool, "checkin", self._remove_conn)
46 event.listen(pool, "close", self._remove_conn)
47 event.listen(pool, "close_detached", self._remove_conn)
48 # note we are keeping "invalidated" here, as those are still
49 # opened connections we would like to roll back
50
51 def _add_conn(self, dbapi_con, con_record, con_proxy):
52 self.dbapi_connections.add(dbapi_con)
53 self.proxy_refs[con_proxy] = True
54
55 def _remove_conn(self, dbapi_conn, *arg):
56 self.dbapi_connections.discard(dbapi_conn)
57
58 def add_engine(self, engine, scope):
59 self.add_pool(engine.pool)
60
61 assert scope in ("class", "global", "function", "fixture")
62 self.testing_engines[scope].add(engine)
63
64 def _safe(self, fn):
65 try:
66 fn()
67 except Exception as e:
68 warnings.warn(
69 "testing_reaper couldn't rollback/close connection: %s" % e
70 )
71
72 def rollback_all(self):
73 for rec in list(self.proxy_refs):
74 if rec is not None and rec.is_valid:
75 self._safe(rec.rollback)
76
77 def checkin_all(self):
78 # run pool.checkin() for all ConnectionFairy instances we have
79 # tracked.
80
81 for rec in list(self.proxy_refs):
82 if rec is not None and rec.is_valid:
83 self.dbapi_connections.discard(rec.dbapi_connection)
84 self._safe(rec._checkin)
85
86 # for fairy refs that were GCed and could not close the connection,
87 # such as asyncio, roll back those remaining connections
88 for con in self.dbapi_connections:
89 self._safe(con.rollback)
90 self.dbapi_connections.clear()
91
92 def close_all(self):
93 self.checkin_all()
94
95 def prepare_for_drop_tables(self, connection):
96 # don't do aggressive checks for third party test suites
97 if not config.bootstrapped_as_sqlalchemy:
98 return
99
100 from . import provision
101
102 provision.prepare_for_drop_tables(connection.engine.url, connection)
103
104 def _drop_testing_engines(self, scope):
105 eng = self.testing_engines[scope]
106 for rec in list(eng):
107 for proxy_ref in list(self.proxy_refs):
108 if proxy_ref is not None and proxy_ref.is_valid:
109 if (
110 proxy_ref._pool is not None
111 and proxy_ref._pool is rec.pool
112 ):
113 self._safe(proxy_ref._checkin)
114
115 if hasattr(rec, "sync_engine"):
116 await_only(rec.dispose())
117 else:
118 rec.dispose()
119
120 eng.clear()
121
122 def _dispose_testing_engines(self, scope):
123 eng = self.testing_engines[scope]
124 for rec in list(eng):
125 if hasattr(rec, "sync_engine"):
126 await_only(rec.dispose())
127 else:
128 rec.dispose()
129
130 def after_test(self):
131 self._drop_testing_engines("function")
132
133 def after_test_outside_fixtures(self, test):
134 # don't do aggressive checks for third party test suites
135 if not config.bootstrapped_as_sqlalchemy:
136 return
137
138 if test.__class__.__leave_connections_for_teardown__:
139 return
140
141 self.checkin_all()
142
143 # on PostgreSQL, this will test for any "idle in transaction"
144 # connections. useful to identify tests with unusual patterns
145 # that can't be cleaned up correctly.
146 from . import provision
147
148 with config.db.connect() as conn:
149 provision.prepare_for_drop_tables(conn.engine.url, conn)
150
151 def stop_test_class_inside_fixtures(self):
152 self.checkin_all()
153 self._drop_testing_engines("function")
154 self._drop_testing_engines("class")
155
156 def stop_test_class_outside_fixtures(self):
157 # ensure no refs to checked out connections at all.
158
159 if pool.base._strong_ref_connection_records:
160 gc_collect()
161
162 if pool.base._strong_ref_connection_records:
163 ln = len(pool.base._strong_ref_connection_records)
164 pool.base._strong_ref_connection_records.clear()
165
166 if ln > 2:
167 # allow two connections to linger, as on loaded down
168 # CI hardware there seem to be occasional GC lapses that
169 # are not easily preventable
170 assert (
171 False
172 ), "%d connection recs not cleared after test suite" % (ln)
173 if config.options and config.options.low_connections:
174 # for suites running with --low-connections, dispose the "global"
175 # engines to disconnect everything before making a testing engine
176 self._dispose_testing_engines("global")
177
178 def final_cleanup(self):
179 self.checkin_all()
180 for scope in self.testing_engines:
181 self._drop_testing_engines(scope)
182
183 def assert_all_closed(self):
184 for rec in self.proxy_refs:
185 if rec.is_valid:
186 assert False
187
188
189testing_reaper = ConnectionKiller()
190
191
192@decorator
193def assert_conns_closed(fn, *args, **kw):
194 try:
195 fn(*args, **kw)
196 finally:
197 testing_reaper.assert_all_closed()
198
199
200@decorator
201def rollback_open_connections(fn, *args, **kw):
202 """Decorator that rolls back all open connections after fn execution."""
203
204 try:
205 fn(*args, **kw)
206 finally:
207 testing_reaper.rollback_all()
208
209
210@decorator
211def close_first(fn, *args, **kw):
212 """Decorator that closes all connections before fn execution."""
213
214 testing_reaper.checkin_all()
215 fn(*args, **kw)
216
217
218@decorator
219def close_open_connections(fn, *args, **kw):
220 """Decorator that closes all connections after fn execution."""
221 try:
222 fn(*args, **kw)
223 finally:
224 testing_reaper.checkin_all()
225
226
227def all_dialects(exclude=None):
228 import sqlalchemy.dialects as d
229
230 for name in d.__all__:
231 # TEMPORARY
232 if exclude and name in exclude:
233 continue
234 mod = getattr(d, name, None)
235 if not mod:
236 mod = getattr(
237 __import__("sqlalchemy.dialects.%s" % name).dialects, name
238 )
239 yield mod.dialect()
240
241
242class ReconnectFixture:
243 def __init__(self, dbapi):
244 self.dbapi = dbapi
245 self.connections = []
246 self.is_stopped = False
247
248 def __getattr__(self, key):
249 return getattr(self.dbapi, key)
250
251 def connect(self, *args, **kwargs):
252 conn = self.dbapi.connect(*args, **kwargs)
253 if self.is_stopped:
254 self._safe(conn.close)
255 curs = conn.cursor() # should fail on Oracle etc.
256 # should fail for everything that didn't fail
257 # above, connection is closed
258 curs.execute("select 1")
259 assert False, "simulated connect failure didn't work"
260 else:
261 self.connections.append(conn)
262 return conn
263
264 def _safe(self, fn):
265 try:
266 fn()
267 except Exception as e:
268 warnings.warn("ReconnectFixture couldn't close connection: %s" % e)
269
270 def shutdown(self, stop=False):
271 # TODO: this doesn't cover all cases
272 # as nicely as we'd like, namely MySQLdb.
273 # would need to implement R. Brewer's
274 # proxy server idea to get better
275 # coverage.
276 self.is_stopped = stop
277 for c in list(self.connections):
278 self._safe(c.close)
279 self.connections = []
280
281 def restart(self):
282 self.is_stopped = False
283
284
285def reconnecting_engine(url=None, options=None):
286 url = url or config.db.url
287 dbapi = config.db.dialect.dbapi
288 if not options:
289 options = {}
290 options["module"] = ReconnectFixture(dbapi)
291 engine = testing_engine(url, options)
292 _dispose = engine.dispose
293
294 def dispose():
295 engine.dialect.dbapi.shutdown()
296 engine.dialect.dbapi.is_stopped = False
297 _dispose()
298
299 engine.test_shutdown = engine.dialect.dbapi.shutdown
300 engine.test_restart = engine.dialect.dbapi.restart
301 engine.dispose = dispose
302 return engine
303
304
305@typing.overload
306def testing_engine(
307 url: Optional[URL] = ...,
308 options: Optional[Dict[str, Any]] = ...,
309 *,
310 asyncio: Literal[False],
311) -> Engine: ...
312
313
314@typing.overload
315def testing_engine(
316 url: Optional[URL] = ...,
317 options: Optional[Dict[str, Any]] = ...,
318 *,
319 asyncio: Literal[True],
320) -> AsyncEngine: ...
321
322
323def testing_engine(
324 url: Optional[URL] = None,
325 options: Optional[Dict[str, Any]] = None,
326 *,
327 asyncio: bool = False,
328) -> Union[Engine, AsyncEngine]:
329
330 if asyncio:
331 from sqlalchemy.ext.asyncio import (
332 create_async_engine as create_engine,
333 )
334 else:
335 from sqlalchemy import create_engine
336 from sqlalchemy.engine.url import make_url
337
338 url = make_url(url if url else config.db.url)
339
340 if not options:
341 options = {}
342
343 use_options = {}
344
345 for opt_dict in (config.db_opts, options):
346 if not opt_dict:
347 continue
348 use_options.update(
349 {
350 opt: value
351 for opt, value in opt_dict.items()
352 if opt not in ("scope", "use_reaper")
353 and not opt.startswith("sqlite_")
354 }
355 )
356
357 engine = create_engine(url, **use_options)
358
359 if config.options and config.options.low_connections:
360 # for suites running with --low-connections, dispose the "global"
361 # engines to disconnect everything before making a testing engine
362 testing_reaper._dispose_testing_engines("global")
363
364 scope = options.get("scope", "function")
365 if scope == "global":
366 if asyncio:
367 engine.sync_engine._has_events = True
368 else:
369 engine._has_events = (
370 True # enable event blocks, helps with profiling
371 )
372
373 from . import provision
374
375 provision.post_configure_testing_engine(engine.url, engine, options, scope)
376
377 # post_configure_testing_engine may have modified the options dictionary
378 # in place; consume additional post arguments afterwards
379
380 use_reaper = options.get("use_reaper", True)
381 if use_reaper:
382 testing_reaper.add_engine(engine, scope)
383
384 if (
385 isinstance(engine.pool, pool.QueuePool)
386 and "pool" not in options
387 and "pool_timeout" not in options
388 and "max_overflow" not in options
389 ):
390 engine.pool._timeout = 0
391 engine.pool._max_overflow = 0
392
393 return engine
394
395
396def mock_engine(dialect_name=None):
397 """Provides a mocking engine based on the current testing.db.
398
399 This is normally used to test DDL generation flow as emitted
400 by an Engine.
401
402 It should not be used in other cases, as assert_compile() and
403 assert_sql_execution() are much better choices with fewer
404 moving parts.
405
406 """
407
408 from sqlalchemy import create_mock_engine
409
410 if not dialect_name:
411 dialect_name = config.db.name
412
413 buffer = []
414
415 def executor(sql, *a, **kw):
416 buffer.append(sql)
417
418 def assert_sql(stmts):
419 recv = [re.sub(r"[\n\t]", "", str(s)) for s in buffer]
420 assert recv == stmts, recv
421
422 def print_sql():
423 d = engine.dialect
424 return "\n".join(str(s.compile(dialect=d)) for s in engine.mock)
425
426 engine = create_mock_engine(dialect_name + "://", executor)
427 assert not hasattr(engine, "mock")
428 engine.mock = buffer
429 engine.assert_sql = assert_sql
430 engine.print_sql = print_sql
431 return engine
432
433
434class DBAPIProxyCursor:
435 """Proxy a DBAPI cursor.
436
437 Tests can provide subclasses of this to intercept
438 DBAPI-level cursor operations.
439
440 """
441
442 def __init__(self, engine, conn, *args, **kwargs):
443 self.engine = engine
444 self.connection = conn
445 self.cursor = conn.cursor(*args, **kwargs)
446
447 def execute(self, stmt, parameters=None, **kw):
448 if parameters:
449 return self.cursor.execute(stmt, parameters, **kw)
450 else:
451 return self.cursor.execute(stmt, **kw)
452
453 def executemany(self, stmt, params, **kw):
454 return self.cursor.executemany(stmt, params, **kw)
455
456 def __iter__(self):
457 return iter(self.cursor)
458
459 def __getattr__(self, key):
460 return getattr(self.cursor, key)
461
462
463class DBAPIProxyConnection:
464 """Proxy a DBAPI connection.
465
466 Tests can provide subclasses of this to intercept
467 DBAPI-level connection operations.
468
469 """
470
471 def __init__(self, engine, conn, cursor_cls):
472 self.conn = conn
473 self.engine = engine
474 self.cursor_cls = cursor_cls
475
476 def cursor(self, *args, **kwargs):
477 return self.cursor_cls(self.engine, self.conn, *args, **kwargs)
478
479 def close(self):
480 self.conn.close()
481
482 def __getattr__(self, key):
483 return getattr(self.conn, key)
484 