Team Ai
Datasetpublic

codekingpro/portable-devtools

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