codekingpro/portable-devtools
115k
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 