Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
plugin_base.py829 linesDownload Raw Back to plugin
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 
codekingpro/portable-devtools · Team Ai