Team Ai
Datasetpublic

codekingpro/portable-devtools

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