Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
engines.py484 linesDownload Raw Back to testing
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 
codekingpro/portable-devtools · Team Ai