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