Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
provision.py604 linesDownload Raw Back to testing
1# testing/provision.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
9from __future__ import annotations
10
11import collections
12import contextlib
13import logging
14
15from . import config
16from . import engines
17from . import util
18from .. import exc
19from .. import inspect
20from ..engine import Connection
21from ..engine import Engine
22from ..engine import url as sa_url
23from ..schema import sort_tables_and_constraints
24from ..sql import ddl
25from ..sql import schema
26from ..util import decorator
27
28
29log = logging.getLogger(__name__)
30
31FOLLOWER_IDENT = None
32
33
34class register:
35    def __init__(self, decorator=None):
36        self.fns = {}
37        self.decorator = decorator
38
39    @classmethod
40    def init(cls, fn):
41        return register().for_db("*")(fn)
42
43    @classmethod
44    def init_decorator(cls, decorator):
45        return register(decorator).for_db("*")
46
47    def for_db(self, *dbnames):
48        def decorate(fn):
49            if self.decorator:
50                fn = self.decorator(fn)
51            for dbname in dbnames:
52                self.fns[dbname] = fn
53            return self
54
55        return decorate
56
57    def call_original(self, cfg, *arg, **kw):
58        return self.fns["*"](cfg, *arg, **kw)
59
60    def __call__(self, cfg, *arg, **kw):
61        if isinstance(cfg, str):
62            url = sa_url.make_url(cfg)
63        elif isinstance(cfg, sa_url.URL):
64            url = cfg
65        elif isinstance(cfg, (Engine, Connection)):
66            url = cfg.engine.url
67        else:
68            url = cfg.db.url
69        backend = url.get_backend_name()
70        if backend in self.fns:
71            return self.fns[backend](cfg, *arg, **kw)
72        else:
73            return self.fns["*"](cfg, *arg, **kw)
74
75
76def create_follower_db(follower_ident):
77    for cfg in _configs_for_db_operation():
78        log.info("CREATE database %s, URI %r", follower_ident, cfg.db.url)
79        create_db(cfg, cfg.db, follower_ident)
80
81
82def setup_config(db_url, options, file_config, follower_ident):
83    # load the dialect, which should also have it set up its provision
84    # hooks
85
86    dialect = sa_url.make_url(db_url).get_dialect()
87
88    dialect.load_provisioning()
89
90    if follower_ident:
91        db_url = follower_url_from_main(db_url, follower_ident)
92    db_opts = {}
93    update_db_opts(db_url, db_opts, options)
94    db_opts["scope"] = "global"
95    eng = engines.testing_engine(db_url, db_opts)
96
97    post_configure_engine(db_url, eng, follower_ident)
98
99    eng.connect().close()
100
101    cfg = config.Config.register(eng, db_opts, options, file_config)
102
103    # a symbolic name that tests can use if they need to disambiguate
104    # names across databases
105    if follower_ident:
106        config.ident = follower_ident
107
108    if follower_ident:
109        configure_follower(cfg, follower_ident)
110    return cfg
111
112
113def drop_follower_db(follower_ident):
114    for cfg in _configs_for_db_operation():
115        log.info("DROP database %s, URI %r", follower_ident, cfg.db.url)
116        drop_db(cfg, cfg.db, follower_ident)
117
118
119def generate_db_urls(db_urls, extra_drivers):
120    """Generate a set of URLs to test given configured URLs plus additional
121    driver names.
122
123    Given:
124
125    .. sourcecode:: text
126
127        --dburi postgresql://db1  \
128        --dburi postgresql://db2  \
129        --dburi postgresql://db2  \
130        --dbdriver=psycopg2 --dbdriver=asyncpg?async_fallback=true
131
132    Noting that the default postgresql driver is psycopg2,  the output
133    would be:
134
135    .. sourcecode:: text
136
137        postgresql+psycopg2://db1
138        postgresql+asyncpg://db1
139        postgresql+psycopg2://db2
140        postgresql+psycopg2://db3
141
142    That is, for the driver in a --dburi, we want to keep that and use that
143    driver for each URL it's part of .   For a driver that is only
144    in --dbdrivers, we want to use it just once for one of the URLs.
145    for a driver that is both coming from --dburi as well as --dbdrivers,
146    we want to keep it in that dburi.
147
148    Driver specific query options can be specified by added them to the
149    driver name. For example, to enable the async fallback option for
150    asyncpg::
151
152    .. sourcecode:: text
153
154        --dburi postgresql://db1  \
155        --dbdriver=asyncpg?async_fallback=true
156
157    """
158    urls = set()
159
160    backend_to_driver_we_already_have = collections.defaultdict(set)
161
162    urls_plus_dialects = [
163        (url_obj, url_obj.get_dialect())
164        for url_obj in [sa_url.make_url(db_url) for db_url in db_urls]
165    ]
166
167    for url_obj, dialect in urls_plus_dialects:
168        # use get_driver_name instead of dialect.driver to account for
169        # "_async" virtual drivers like oracledb and psycopg
170        driver_name = url_obj.get_driver_name()
171        backend_to_driver_we_already_have[dialect.name].add(driver_name)
172
173    backend_to_driver_we_need = {}
174
175    for url_obj, dialect in urls_plus_dialects:
176        backend = dialect.name
177        dialect.load_provisioning()
178
179        if backend not in backend_to_driver_we_need:
180            backend_to_driver_we_need[backend] = extra_per_backend = set(
181                extra_drivers
182            ).difference(backend_to_driver_we_already_have[backend])
183        else:
184            extra_per_backend = backend_to_driver_we_need[backend]
185
186        for driver_url in _generate_driver_urls(url_obj, extra_per_backend):
187            if driver_url in urls:
188                continue
189            urls.add(driver_url)
190            yield driver_url
191
192
193def _generate_driver_urls(url, extra_drivers):
194    main_driver = url.get_driver_name()
195    extra_drivers.discard(main_driver)
196
197    url = generate_driver_url(url, main_driver, "")
198    yield url
199
200    for drv in list(extra_drivers):
201        if "?" in drv:
202            driver_only, query_str = drv.split("?", 1)
203
204        else:
205            driver_only = drv
206            query_str = None
207
208        new_url = generate_driver_url(url, driver_only, query_str)
209        if new_url:
210            extra_drivers.remove(drv)
211
212            yield new_url
213
214
215@register.init
216def is_preferred_driver(cfg, engine):
217    """Return True if the engine's URL is on the "default" driver, or
218    more generally the "preferred" driver to use for tests.
219
220    Backends can override this to make a different driver the "prefeferred"
221    driver that's not the default.
222
223    """
224    return (
225        engine.url._get_entrypoint()
226        is engine.url.set(
227            drivername=engine.url.get_backend_name()
228        )._get_entrypoint()
229    )
230
231
232@register.init
233def generate_driver_url(url, driver, query_str):
234    backend = url.get_backend_name()
235
236    new_url = url.set(
237        drivername="%s+%s" % (backend, driver),
238    )
239    if query_str:
240        new_url = new_url.update_query_string(query_str)
241
242    try:
243        new_url.get_dialect()
244    except exc.NoSuchModuleError:
245        return None
246    else:
247        return new_url
248
249
250def _configs_for_db_operation():
251    hosts = set()
252
253    for cfg in config.Config.all_configs():
254        cfg.db.dispose()
255
256    for cfg in config.Config.all_configs():
257        url = cfg.db.url
258        backend = url.get_backend_name()
259        host_conf = (backend, url.username, url.host, url.database)
260
261        if host_conf not in hosts:
262            yield cfg
263            hosts.add(host_conf)
264
265    for cfg in config.Config.all_configs():
266        cfg.db.dispose()
267
268
269@register.init
270def drop_all_schema_objects_pre_tables(cfg, eng):
271    pass
272
273
274@register.init
275def drop_all_schema_objects_post_tables(cfg, eng):
276    pass
277
278
279def drop_all_schema_objects(cfg, eng):
280    drop_all_schema_objects_pre_tables(cfg, eng)
281
282    drop_views(cfg, eng)
283
284    if config.requirements.materialized_views.enabled:
285        drop_materialized_views(cfg, eng)
286
287    inspector = inspect(eng)
288
289    consider_schemas = (None,)
290    if config.requirements.schemas.enabled_for_config(cfg):
291        consider_schemas += (cfg.test_schema, cfg.test_schema_2)
292    util.drop_all_tables(eng, inspector, consider_schemas=consider_schemas)
293
294    drop_all_schema_objects_post_tables(cfg, eng)
295
296    if config.requirements.sequences.enabled_for_config(cfg):
297        with eng.begin() as conn:
298            for seq in inspector.get_sequence_names():
299                conn.execute(ddl.DropSequence(schema.Sequence(seq)))
300            if config.requirements.schemas.enabled_for_config(cfg):
301                for schema_name in [cfg.test_schema, cfg.test_schema_2]:
302                    for seq in inspector.get_sequence_names(
303                        schema=schema_name
304                    ):
305                        conn.execute(
306                            ddl.DropSequence(
307                                schema.Sequence(seq, schema=schema_name)
308                            )
309                        )
310
311
312def drop_views(cfg, eng):
313    inspector = inspect(eng)
314
315    try:
316        view_names = inspector.get_view_names()
317    except NotImplementedError:
318        pass
319    else:
320        with eng.begin() as conn:
321            for vname in view_names:
322                conn.execute(
323                    ddl._DropView(schema.Table(vname, schema.MetaData()))
324                )
325
326    if config.requirements.schemas.enabled_for_config(cfg):
327        try:
328            view_names = inspector.get_view_names(schema=cfg.test_schema)
329        except NotImplementedError:
330            pass
331        else:
332            with eng.begin() as conn:
333                for vname in view_names:
334                    conn.execute(
335                        ddl._DropView(
336                            schema.Table(
337                                vname,
338                                schema.MetaData(),
339                                schema=cfg.test_schema,
340                            )
341                        )
342                    )
343
344
345def drop_materialized_views(cfg, eng):
346    inspector = inspect(eng)
347
348    mview_names = inspector.get_materialized_view_names()
349
350    with eng.begin() as conn:
351        for vname in mview_names:
352            conn.exec_driver_sql(f"DROP MATERIALIZED VIEW {vname}")
353
354    if config.requirements.schemas.enabled_for_config(cfg):
355        mview_names = inspector.get_materialized_view_names(
356            schema=cfg.test_schema
357        )
358        with eng.begin() as conn:
359            for vname in mview_names:
360                conn.exec_driver_sql(
361                    f"DROP MATERIALIZED VIEW {cfg.test_schema}.{vname}"
362                )
363
364
365@register.init
366def create_db(cfg, eng, ident):
367    """Dynamically create a database for testing.
368
369    Used when a test run will employ multiple processes, e.g., when run
370    via `tox` or `pytest -n4`.
371    """
372    raise NotImplementedError(
373        "no DB creation routine for cfg: %s" % (eng.url,)
374    )
375
376
377@register.init
378def drop_db(cfg, eng, ident):
379    """Drop a database that we dynamically created for testing."""
380    raise NotImplementedError("no DB drop routine for cfg: %s" % (eng.url,))
381
382
383def _adapt_update_db_opts(fn):
384    insp = util.inspect_getfullargspec(fn)
385    if len(insp.args) == 3:
386        return fn
387    else:
388        return lambda db_url, db_opts, _options: fn(db_url, db_opts)
389
390
391@register.init_decorator(_adapt_update_db_opts)
392def update_db_opts(db_url, db_opts, options):
393    """Set database options (db_opts) for a test database that we created."""
394
395
396@register.init
397def post_configure_engine(url, engine, follower_ident):
398    """Perform extra steps after configuring the main engine for testing.
399
400    (For the internal dialects, currently only used by sqlite, oracle, mssql)
401    """
402
403
404@register.init
405def post_configure_testing_engine(url, engine, options, scope):
406    """perform extra steps after configuring any engine within the
407    testing_engine() function.
408
409    this includes the main engine as well as most ad-hoc testing engines.
410
411    steps here should not get in the way of test cases that are looking
412    for events, etc.
413
414    """
415
416
417@register.init
418def follower_url_from_main(url, ident):
419    """Create a connection URL for a dynamically-created test database.
420
421    :param url: the connection URL specified when the test run was invoked
422    :param ident: the pytest-xdist "worker identifier" to be used as the
423                  database name
424    """
425    url = sa_url.make_url(url)
426    return url.set(database=ident)
427
428
429@register.init
430def configure_follower(cfg, ident):
431    """Create dialect-specific config settings for a follower database."""
432    pass
433
434
435@register.init
436def run_reap_dbs(url, ident):
437    """Remove databases that were created during the test process, after the
438    process has ended.
439
440    This is an optional step that is invoked for certain backends that do not
441    reliably release locks on the database as long as a process is still in
442    use. For the internal dialects, this is currently only necessary for
443    mssql and oracle.
444    """
445
446
447def reap_dbs(idents_file):
448    log.info("Reaping databases...")
449
450    urls = collections.defaultdict(set)
451    idents = collections.defaultdict(set)
452    dialects = {}
453
454    with open(idents_file) as file_:
455        for line in file_:
456            line = line.strip()
457            db_name, db_url = line.split(" ")
458            url_obj = sa_url.make_url(db_url)
459            if db_name not in dialects:
460                dialects[db_name] = url_obj.get_dialect()
461                dialects[db_name].load_provisioning()
462            url_key = (url_obj.get_backend_name(), url_obj.host)
463            urls[url_key].add(db_url)
464            idents[url_key].add(db_name)
465
466    for url_key in urls:
467        url = list(urls[url_key])[0]
468        ident = idents[url_key]
469        run_reap_dbs(url, ident)
470
471
472@register.init
473def temp_table_keyword_args(cfg, eng):
474    """Specify keyword arguments for creating a temporary Table.
475
476    Dialect-specific implementations of this method will return the
477    kwargs that are passed to the Table method when creating a temporary
478    table for testing, e.g., in the define_temp_tables method of the
479    ComponentReflectionTest class in suite/test_reflection.py
480    """
481    raise NotImplementedError(
482        "no temp table keyword args routine for cfg: %s" % (eng.url,)
483    )
484
485
486@register.init
487def prepare_for_drop_tables(config, connection):
488    pass
489
490
491@register.init
492def stop_test_class_outside_fixtures(config, db, testcls):
493    pass
494
495
496@register.init
497def get_temp_table_name(cfg, eng, base_name):
498    """Specify table name for creating a temporary Table.
499
500    Dialect-specific implementations of this method will return the
501    name to use when creating a temporary table for testing,
502    e.g., in the define_temp_tables method of the
503    ComponentReflectionTest class in suite/test_reflection.py
504
505    Default to just the base name since that's what most dialects will
506    use. The mssql dialect's implementation will need a "#" prepended.
507    """
508    return base_name
509
510
511@register.init
512def set_default_schema_on_connection(cfg, dbapi_connection, schema_name):
513    raise NotImplementedError(
514        "backend does not implement a schema name set function: %s"
515        % (cfg.db.url,)
516    )
517
518
519@register.init
520def upsert(
521    cfg,
522    table,
523    returning,
524    *,
525    set_lambda=None,
526    sort_by_parameter_order=False,
527    index_elements=None,
528):
529    """return the backends insert..on conflict / on dupe etc. construct.
530
531    while we should add a backend-neutral upsert construct as well, such as
532    insert().upsert(), it's important that we continue to test the
533    backend-specific insert() constructs since if we do implement
534    insert().upsert(), that would be using a different codepath for the things
535    we need to test like insertmanyvalues, etc.
536
537    """
538    raise NotImplementedError(
539        f"backend does not include an upsert implementation: {cfg.db.url}"
540    )
541
542
543@register.init
544def normalize_sequence(cfg, sequence):
545    """Normalize sequence parameters for dialect that don't start with 1
546    by default.
547
548    The default implementation does nothing
549    """
550    return sequence
551
552
553@register.init
554def allow_stale_update_impl(cfg):
555    return contextlib.nullcontext()
556
557
558@decorator
559def allow_stale_updates(fn, *arg, **kw):
560    """decorator around a test function that indicates the test will
561    be UPDATING rows that have been read and are now stale.
562
563    This normally doesn't require intervention except for mariadb 12
564    which now raises its own error for that, and we want to turn off
565    that setting just within the scope of the test that needs it
566    to be turned off (i.e. ORM stale version tests)
567
568    """
569    with allow_stale_update_impl(config._current):
570        return fn(*arg, **kw)
571
572
573@register.init
574def delete_from_all_tables(connection, cfg, metadata):
575    """an absolutely foolproof delete from all tables routine.
576
577    dialects should override this to add special instructions like
578    disable constraints etc.
579
580    """
581    savepoints = getattr(cfg.requirements, "savepoints", False)
582    if savepoints:
583        savepoints = savepoints.enabled
584
585    inspector = inspect(connection)
586
587    for table in reversed(
588        [
589            t
590            for (t, fks) in sort_tables_and_constraints(
591                metadata.tables.values()
592            )
593            if t is not None
594            # remember that inspector.get_table_names() is cached,
595            # so this emits SQL once per unique schema name
596            and t.name in inspector.get_table_names(schema=t.schema)
597        ]
598    ):
599        if savepoints:
600            with connection.begin_nested():
601                connection.execute(table.delete())
602        else:
603            connection.execute(table.delete())
604 
codekingpro/portable-devtools · Team Ai