codekingpro/portable-devtools
115k
1# testing/plugin/plugin_base.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
9
10from __future__ import annotations
11
12import abc
13from argparse import Namespace
14import configparser
15import logging
16import os
17from pathlib import Path
18import re
19import sys
20from typing import Any
21
22from sqlalchemy.testing import asyncio
23
24"""Testing extensions.
25
26this module is designed to work as a testing-framework-agnostic library,
27created so that multiple test frameworks can be supported at once
28(mostly so that we can migrate to new ones). The current target
29is pytest.
30
31"""
32
33# flag which indicates we are in the SQLAlchemy testing suite,
34# and not that of Alembic or a third party dialect.
35bootstrapped_as_sqlalchemy = False
36
37log = logging.getLogger("sqlalchemy.testing.plugin_base")
38
39# late imports
40fixtures = None
41engines = None
42exclusions = None
43warnings = None
44profiling = None
45provision = None
46assertions = None
47requirements = None
48config = None
49testing = None
50util = None
51file_config = None
52
53logging = None
54include_tags = set()
55exclude_tags = set()
56options: Namespace = None # type: ignore
57
58
59def setup_options(make_option):
60 make_option(
61 "--log-info",
62 action="callback",
63 type=str,
64 callback=_log,
65 help="turn on info logging for <LOG> (multiple OK)",
66 )
67 make_option(
68 "--log-debug",
69 action="callback",
70 type=str,
71 callback=_log,
72 help="turn on debug logging for <LOG> (multiple OK)",
73 )
74 make_option(
75 "--db",
76 action="append",
77 type=str,
78 dest="db",
79 help="Use prefab database uri. Multiple OK, "
80 "first one is run by default.",
81 )
82 make_option(
83 "--dbs",
84 action="callback",
85 zeroarg_callback=_list_dbs,
86 help="List available prefab dbs",
87 )
88 make_option(
89 "--dburi",
90 action="append",
91 type=str,
92 dest="dburi",
93 help="Database uri. Multiple OK, first one is run by default.",
94 )
95 make_option(
96 "--dbdriver",
97 action="append",
98 type=str,
99 dest="dbdriver",
100 help="Additional database drivers to include in tests. "
101 "These are linked to the existing database URLs by the "
102 "provisioning system.",
103 )
104 make_option(
105 "--dropfirst",
106 action="store_true",
107 dest="dropfirst",
108 help="Drop all tables in the target database first",
109 )
110 make_option(
111 "--disable-asyncio",
112 action="store_true",
113 help="disable test / fixtures / provisioning running in asyncio",
114 )
115 make_option(
116 "--backend-only",
117 action="callback",
118 zeroarg_callback=_set_tag_include("backend"),
119 help=(
120 "Run only tests marked with __backend__ or __sparse_backend__; "
121 "this is now equivalent to the pytest -m backend mark expression"
122 ),
123 )
124 make_option(
125 "--nomemory",
126 action="callback",
127 zeroarg_callback=_set_tag_exclude("memory_intensive"),
128 help="Don't run memory profiling tests; "
129 "this is now equivalent to the pytest -m 'not memory_intensive' "
130 "mark expression",
131 )
132 make_option(
133 "--notimingintensive",
134 action="callback",
135 zeroarg_callback=_set_tag_exclude("timing_intensive"),
136 help="Don't run timing intensive tests; "
137 "this is now equivalent to the pytest -m 'not timing_intensive' "
138 "mark expression",
139 )
140 make_option(
141 "--nomypy",
142 action="callback",
143 zeroarg_callback=_set_tag_exclude("mypy"),
144 help="Don't run mypy typing tests; "
145 "this is now equivalent to the pytest -m 'not mypy' mark expression",
146 )
147 make_option(
148 "--profile-sort",
149 type=str,
150 default="cumulative",
151 dest="profilesort",
152 help="Type of sort for profiling standard output",
153 )
154 make_option(
155 "--profile-dump",
156 type=str,
157 dest="profiledump",
158 help="Filename where a single profile run will be dumped",
159 )
160 make_option(
161 "--low-connections",
162 action="store_true",
163 dest="low_connections",
164 help="Use a low number of distinct connections - "
165 "i.e. for Oracle TNS",
166 )
167 make_option(
168 "--write-idents",
169 type=str,
170 dest="write_idents",
171 help="write out generated follower idents to <file>, "
172 "when -n<num> is used",
173 )
174 make_option(
175 "--requirements",
176 action="callback",
177 type=str,
178 callback=_requirements_opt,
179 help="requirements class for testing, overrides setup.cfg",
180 )
181 make_option(
182 "--include-tag",
183 action="callback",
184 callback=_include_tag,
185 type=str,
186 help="Include tests with tag <tag>; "
187 "legacy, use pytest -m 'tag' instead",
188 )
189 make_option(
190 "--exclude-tag",
191 action="callback",
192 callback=_exclude_tag,
193 type=str,
194 help="Exclude tests with tag <tag>; "
195 "legacy, use pytest -m 'not tag' instead",
196 )
197 make_option(
198 "--write-profiles",
199 action="store_true",
200 dest="write_profiles",
201 default=False,
202 help="Write/update failing profiling data.",
203 )
204 make_option(
205 "--force-write-profiles",
206 action="store_true",
207 dest="force_write_profiles",
208 default=False,
209 help="Unconditionally write/update profiling data.",
210 )
211 make_option(
212 "--dump-pyannotate",
213 type=str,
214 dest="dump_pyannotate",
215 help="Run pyannotate and dump json info to given file",
216 )
217 make_option(
218 "--mypy-extra-test-path",
219 type=str,
220 action="append",
221 default=[],
222 dest="mypy_extra_test_paths",
223 help="Additional test directories to add to the mypy tests. "
224 "This is used only when running mypy tests. Multiple OK",
225 )
226 # db specific options
227 make_option(
228 "--postgresql-templatedb",
229 type=str,
230 help="name of template database to use for PostgreSQL "
231 "CREATE DATABASE (defaults to current database)",
232 )
233 make_option(
234 "--oracledb-thick-mode",
235 action="store_true",
236 help="enables the 'thick mode' when testing with oracle+oracledb",
237 )
238
239
240def configure_follower(follower_ident):
241 """Configure required state for a follower.
242
243 This invokes in the parent process and typically includes
244 database creation.
245
246 """
247 from sqlalchemy.testing import provision
248
249 provision.FOLLOWER_IDENT = follower_ident
250
251
252def memoize_important_follower_config(dict_):
253 """Store important configuration we will need to send to a follower.
254
255 This invokes in the parent process after normal config is set up.
256
257 Hook is currently not used.
258
259 """
260
261
262def restore_important_follower_config(dict_):
263 """Restore important configuration needed by a follower.
264
265 This invokes in the follower process.
266
267 Hook is currently not used.
268
269 """
270
271
272def read_config(root_path):
273 global file_config
274 file_config = configparser.ConfigParser()
275 file_config.read(
276 [str(root_path / "setup.cfg"), str(root_path / "test.cfg")]
277 )
278
279
280def pre_begin(opt):
281 """things to set up early, before coverage might be setup."""
282 global options
283 options = opt
284 for fn in pre_configure:
285 fn(options, file_config)
286
287
288def set_coverage_flag(value):
289 options.has_coverage = value
290
291
292def post_begin():
293 """things to set up later, once we know coverage is running."""
294 # Lazy setup of other options (post coverage)
295 for fn in post_configure:
296 fn(options, file_config)
297
298 # late imports, has to happen after config.
299 global util, fixtures, engines, exclusions, assertions, provision
300 global warnings, profiling, config, testing
301 from sqlalchemy import testing # noqa
302 from sqlalchemy.testing import fixtures, engines, exclusions # noqa
303 from sqlalchemy.testing import assertions, warnings, profiling # noqa
304 from sqlalchemy.testing import config, provision # noqa
305 from sqlalchemy import util # noqa
306
307 warnings.setup_filters()
308
309
310def _log(opt_str, value, parser):
311 global logging
312 if not logging:
313 import logging
314
315 logging.basicConfig()
316
317 if opt_str.endswith("-info"):
318 logging.getLogger(value).setLevel(logging.INFO)
319 elif opt_str.endswith("-debug"):
320 logging.getLogger(value).setLevel(logging.DEBUG)
321
322
323def _list_dbs(*args):
324 if file_config is None:
325 # assume the current working directory is the one containing the
326 # setup file
327 read_config(Path.cwd())
328 print("Available --db options (use --dburi to override)")
329 for macro in sorted(file_config.options("db")):
330 print("%20s\t%s" % (macro, file_config.get("db", macro)))
331 sys.exit(0)
332
333
334def _requirements_opt(opt_str, value, parser):
335 _setup_requirements(value)
336
337
338def _set_tag_include(tag):
339 def _do_include_tag(opt_str, value, parser):
340 _include_tag(opt_str, tag, parser)
341
342 return _do_include_tag
343
344
345def _set_tag_exclude(tag):
346 def _do_exclude_tag(opt_str, value, parser):
347 _exclude_tag(opt_str, tag, parser)
348
349 return _do_exclude_tag
350
351
352def _exclude_tag(opt_str, value, parser):
353 exclude_tags.add(value.replace("-", "_"))
354
355
356def _include_tag(opt_str, value, parser):
357 include_tags.add(value.replace("-", "_"))
358
359
360pre_configure = []
361post_configure = []
362
363
364def pre(fn):
365 pre_configure.append(fn)
366 return fn
367
368
369def post(fn):
370 post_configure.append(fn)
371 return fn
372
373
374@pre
375def _setup_options(opt, file_config):
376 global options
377 options = opt
378
379
380@pre
381def _register_sqlite_numeric_dialect(opt, file_config):
382 from sqlalchemy.dialects import registry
383
384 registry.register(
385 "sqlite.pysqlite_numeric",
386 "sqlalchemy.dialects.sqlite.pysqlite",
387 "_SQLiteDialect_pysqlite_numeric",
388 )
389 registry.register(
390 "sqlite.pysqlite_dollar",
391 "sqlalchemy.dialects.sqlite.pysqlite",
392 "_SQLiteDialect_pysqlite_dollar",
393 )
394
395
396@post
397def __ensure_cext(opt, file_config):
398 if os.environ.get("REQUIRE_SQLALCHEMY_CEXT", "0") == "1":
399 from sqlalchemy.util import has_compiled_ext
400
401 try:
402 has_compiled_ext(raise_=True)
403 except ImportError as err:
404 raise AssertionError(
405 "REQUIRE_SQLALCHEMY_CEXT is set but can't import the "
406 "cython extensions"
407 ) from err
408
409
410@post
411def _init_symbols(options, file_config):
412 from sqlalchemy.testing import config
413
414 config._fixture_functions = _fixture_fn_class()
415
416
417@pre
418def _set_disable_asyncio(opt, file_config):
419 if opt.disable_asyncio:
420 asyncio.ENABLE_ASYNCIO = False
421
422
423@post
424def _engine_uri(options, file_config):
425 from sqlalchemy import testing
426 from sqlalchemy.testing import config
427 from sqlalchemy.testing import provision
428 from sqlalchemy.engine import url as sa_url
429
430 if options.dburi:
431 db_urls = list(options.dburi)
432 else:
433 db_urls = []
434
435 extra_drivers = options.dbdriver or []
436
437 if options.db:
438 for db_token in options.db:
439 for db in re.split(r"[,\s]+", db_token):
440 if db not in file_config.options("db"):
441 raise RuntimeError(
442 "Unknown URI specifier '%s'. "
443 "Specify --dbs for known uris." % db
444 )
445 else:
446 db_urls.append(file_config.get("db", db))
447
448 if not db_urls:
449 db_urls.append(file_config.get("db", "default"))
450
451 config._current = None
452
453 if options.write_idents and provision.FOLLOWER_IDENT:
454 for db_url in [sa_url.make_url(db_url) for db_url in db_urls]:
455 with open(options.write_idents, "a") as file_:
456 file_.write(
457 f"{provision.FOLLOWER_IDENT} "
458 f"{db_url.render_as_string(hide_password=False)}\n"
459 )
460
461 expanded_urls = list(provision.generate_db_urls(db_urls, extra_drivers))
462
463 for db_url in expanded_urls:
464 log.info("Adding database URL: %s", db_url)
465
466 cfg = provision.setup_config(
467 db_url, options, file_config, provision.FOLLOWER_IDENT
468 )
469 if not config._current:
470 cfg.set_as_current(cfg, testing)
471
472
473@post
474def _requirements(options, file_config):
475 requirement_cls = file_config.get("sqla_testing", "requirement_cls")
476 _setup_requirements(requirement_cls)
477
478
479def _setup_requirements(argument):
480 from sqlalchemy.testing import config
481 from sqlalchemy import testing
482
483 modname, clsname = argument.split(":")
484
485 # importlib.import_module() only introduced in 2.7, a little
486 # late
487 mod = __import__(modname)
488 for component in modname.split(".")[1:]:
489 mod = getattr(mod, component)
490 req_cls = getattr(mod, clsname)
491
492 config.requirements = testing.requires = req_cls()
493
494 config.bootstrapped_as_sqlalchemy = bootstrapped_as_sqlalchemy
495
496
497@post
498def _prep_testing_database(options, file_config):
499 from sqlalchemy.testing import config
500
501 if options.dropfirst:
502 from sqlalchemy.testing import provision
503
504 for cfg in config.Config.all_configs():
505 provision.drop_all_schema_objects(cfg, cfg.db)
506
507
508@post
509def _post_setup_options(opt, file_config):
510 from sqlalchemy.testing import config
511
512 config.options = options
513 config.file_config = file_config
514
515
516@post
517def _setup_profiling(options, file_config):
518 from sqlalchemy.testing import profiling
519
520 profiling._profile_stats = profiling.ProfileStatsFile(
521 file_config.get("sqla_testing", "profile_file"),
522 sort=options.profilesort,
523 dump=options.profiledump,
524 )
525
526
527def want_class(name, cls):
528 if not issubclass(cls, fixtures.TestBase):
529 return False
530 elif name.startswith("_"):
531 return False
532 else:
533 return True
534
535
536def want_method(cls, fn):
537 if not fn.__name__.startswith("test_"):
538 return False
539 elif fn.__module__ is None:
540 return False
541 else:
542 return True
543
544
545def generate_sub_tests(cls, module, markers):
546 if (
547 "backend" in markers
548 or "sparse_backend" in markers
549 or "sparse_driver_backend" in markers
550 ):
551 sparse = "sparse_backend" in markers
552 sparse_driver = "sparse_driver_backend" in markers
553 for cfg in _possible_configs_for_cls(
554 cls, sparse=sparse, sparse_driver=sparse_driver
555 ):
556 orig_name = cls.__name__
557
558 # we can have special chars in these names except for the
559 # pytest junit plugin, which is tripped up by the brackets
560 # and periods, so sanitize
561
562 alpha_name = re.sub(r"[_\[\]\.]+", "_", cfg.name)
563 alpha_name = re.sub(r"_+$", "", alpha_name)
564 name = "%s_%s" % (cls.__name__, alpha_name)
565 subcls = type(
566 name,
567 (cls,),
568 {"_sa_orig_cls_name": orig_name, "__only_on_config__": cfg},
569 )
570 setattr(module, name, subcls)
571 yield subcls
572 else:
573 yield cls
574
575
576def start_test_class_outside_fixtures(cls):
577 _do_skips(cls)
578 _setup_engine(cls)
579
580
581def stop_test_class(cls):
582 # close sessions, immediate connections, etc.
583 fixtures.stop_test_class_inside_fixtures(cls)
584
585 # close outstanding connection pool connections, dispose of
586 # additional engines
587 engines.testing_reaper.stop_test_class_inside_fixtures()
588
589
590def stop_test_class_outside_fixtures(cls):
591 provision.stop_test_class_outside_fixtures(config, config.db, cls)
592 engines.testing_reaper.stop_test_class_outside_fixtures()
593 try:
594 if not options.low_connections:
595 assertions.global_cleanup_assertions()
596 finally:
597 _restore_engine()
598
599
600def _restore_engine():
601 if config._current:
602 config._current.reset(testing)
603
604
605def final_process_cleanup():
606 engines.testing_reaper.final_cleanup()
607 assertions.global_cleanup_assertions()
608 _restore_engine()
609
610
611def _setup_engine(cls):
612 if getattr(cls, "__engine_options__", None):
613 opts = dict(cls.__engine_options__)
614 opts["scope"] = "class"
615 eng = engines.testing_engine(options=opts)
616 config._current.push_engine(eng, testing)
617
618
619def before_test(test, test_module_name, test_class, test_name):
620 # format looks like:
621 # "test.aaa_profiling.test_compiler.CompileTest.test_update_whereclause"
622
623 name = getattr(test_class, "_sa_orig_cls_name", test_class.__name__)
624
625 id_ = "%s.%s.%s" % (test_module_name, name, test_name)
626
627 profiling._start_current_test(id_)
628
629
630def after_test(test):
631 fixtures.after_test()
632 engines.testing_reaper.after_test()
633
634
635def after_test_fixtures(test):
636 engines.testing_reaper.after_test_outside_fixtures(test)
637
638
639def _possible_configs_for_cls(
640 cls, reasons=None, sparse=False, sparse_driver=False
641):
642 all_configs = set(config.Config.all_configs())
643
644 if cls.__unsupported_on__:
645 spec = exclusions.db_spec(*cls.__unsupported_on__)
646 for config_obj in list(all_configs):
647 if spec(config_obj):
648 all_configs.remove(config_obj)
649
650 if getattr(cls, "__only_on__", None):
651 spec = exclusions.db_spec(*util.to_list(cls.__only_on__))
652 for config_obj in list(all_configs):
653 if not spec(config_obj):
654 all_configs.remove(config_obj)
655
656 if getattr(cls, "__only_on_config__", None):
657 all_configs.intersection_update([cls.__only_on_config__])
658
659 if hasattr(cls, "__requires__"):
660 requirements = config.requirements
661 for config_obj in list(all_configs):
662 for requirement in cls.__requires__:
663 check = getattr(requirements, requirement)
664
665 skip_reasons = check.matching_config_reasons(config_obj)
666 if skip_reasons:
667 all_configs.remove(config_obj)
668 if reasons is not None:
669 reasons.extend(skip_reasons)
670 break
671
672 warnings = check.matching_warnings(config_obj)
673 if warnings:
674 cls.__warnings__ = getattr(
675 cls, "__warnings__", ()
676 ) + tuple(warnings)
677
678 if hasattr(cls, "__prefer_requires__"):
679 non_preferred = set()
680 requirements = config.requirements
681 for config_obj in list(all_configs):
682 for requirement in cls.__prefer_requires__:
683 check = getattr(requirements, requirement)
684
685 if not check.enabled_for_config(config_obj):
686 non_preferred.add(config_obj)
687 if all_configs.difference(non_preferred):
688 all_configs.difference_update(non_preferred)
689
690 if sparse:
691 # pick only one config from each base dialect
692 # sorted so we get the same backend each time selecting the highest
693 # server version info.
694 per_dialect = {}
695
696 sorted_all_configs = reversed(
697 sorted(
698 all_configs,
699 key=lambda cfg: (
700 "z" if cfg.is_default_dialect else "a",
701 cfg.db.name,
702 cfg.db.driver,
703 cfg.db.dialect.server_version_info,
704 cfg.db.dialect.is_async,
705 ),
706 )
707 )
708
709 for cfg in sorted_all_configs:
710 db = cfg.db.name
711 if db not in per_dialect:
712 per_dialect[db] = cfg
713 return per_dialect.values()
714 elif sparse_driver:
715 # a more liberal form of "sparse" that will select for one driver,
716 # but still return for multiple database servers
717
718 dbs = {}
719
720 sorted_all_configs = list(
721 reversed(
722 sorted(
723 all_configs,
724 key=lambda cfg: (
725 cfg.db.name,
726 cfg.db.driver,
727 cfg.db.dialect.server_version_info,
728 cfg.db.dialect.is_async,
729 ),
730 )
731 )
732 )
733
734 for cfg in sorted_all_configs:
735 key = (cfg.db.name, cfg.db.dialect.server_version_info)
736 if key in dbs and dbs[key].is_default_dialect:
737 continue
738 else:
739 dbs[key] = cfg
740
741 chosen_cfgs = set(dbs.values())
742 return [cfg for cfg in sorted_all_configs if cfg in chosen_cfgs]
743
744 return all_configs
745
746
747def _do_skips(cls):
748 reasons = []
749 all_configs = _possible_configs_for_cls(cls, reasons)
750
751 if getattr(cls, "__skip_if__", False):
752 for c in getattr(cls, "__skip_if__"):
753 if c():
754 config.skip_test(
755 "'%s' skipped by %s" % (cls.__name__, c.__name__)
756 )
757
758 if not all_configs:
759 msg = "'%s.%s' unsupported on any DB implementation %s%s" % (
760 cls.__module__,
761 cls.__name__,
762 ", ".join(
763 "'%s(%s)+%s'"
764 % (
765 config_obj.db.name,
766 ".".join(
767 str(dig)
768 for dig in exclusions._server_version(config_obj.db)
769 ),
770 config_obj.db.driver,
771 )
772 for config_obj in config.Config.all_configs()
773 ),
774 ", ".join(reasons),
775 )
776 config.skip_test(msg)
777 elif hasattr(cls, "__prefer_backends__"):
778 non_preferred = set()
779 spec = exclusions.db_spec(*util.to_list(cls.__prefer_backends__))
780 for config_obj in all_configs:
781 if not spec(config_obj):
782 non_preferred.add(config_obj)
783 if all_configs.difference(non_preferred):
784 all_configs.difference_update(non_preferred)
785
786 if config._current not in all_configs:
787 _setup_config(all_configs.pop(), cls)
788
789
790def _setup_config(config_obj, ctx):
791 config._current.push(config_obj, testing)
792
793
794class FixtureFunctions(abc.ABC):
795 @abc.abstractmethod
796 def skip_test_exception(self, *arg, **kw):
797 raise NotImplementedError()
798
799 @abc.abstractmethod
800 def combinations(self, *args, **kw):
801 raise NotImplementedError()
802
803 @abc.abstractmethod
804 def param_ident(self, *args, **kw):
805 raise NotImplementedError()
806
807 @abc.abstractmethod
808 def fixture(self, *arg, **kw):
809 raise NotImplementedError()
810
811 def get_current_test_name(self):
812 raise NotImplementedError()
813
814 @abc.abstractmethod
815 def mark_base_test_class(self) -> Any:
816 raise NotImplementedError()
817
818 @abc.abstractproperty
819 def add_to_marker(self):
820 raise NotImplementedError()
821
822
823_fixture_fn_class = None
824
825
826def set_fixture_functions(fixture_fn_class):
827 global _fixture_fn_class
828 _fixture_fn_class = fixture_fn_class
829 