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