codekingpro/portable-devtools
115k
1# testing/plugin/pytestplugin.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 argparse
12import collections
13from functools import update_wrapper
14import inspect
15import itertools
16import operator
17import os
18import re
19import sys
20from typing import TYPE_CHECKING
21import uuid
22
23import pytest
24
25try:
26 # installed by bootstrap.py
27 if not TYPE_CHECKING:
28 import sqla_plugin_base as plugin_base
29except ImportError:
30 # assume we're a package, use traditional import
31 from . import plugin_base
32
33
34def pytest_addoption(parser):
35 group = parser.getgroup("sqlalchemy")
36
37 def make_option(name, **kw):
38 callback_ = kw.pop("callback", None)
39 if callback_:
40
41 class CallableAction(argparse.Action):
42 def __call__(
43 self, parser, namespace, values, option_string=None
44 ):
45 callback_(option_string, values, parser)
46
47 kw["action"] = CallableAction
48
49 zeroarg_callback = kw.pop("zeroarg_callback", None)
50 if zeroarg_callback:
51
52 class CallableAction(argparse.Action):
53 def __init__(
54 self,
55 option_strings,
56 dest,
57 default=False,
58 required=False,
59 help=None, # noqa
60 ):
61 super().__init__(
62 option_strings=option_strings,
63 dest=dest,
64 nargs=0,
65 const=True,
66 default=default,
67 required=required,
68 help=help,
69 )
70
71 def __call__(
72 self, parser, namespace, values, option_string=None
73 ):
74 zeroarg_callback(option_string, values, parser)
75
76 kw["action"] = CallableAction
77
78 group.addoption(name, **kw)
79
80 plugin_base.setup_options(make_option)
81
82
83def pytest_configure(config: pytest.Config):
84 plugin_base.read_config(config.rootpath)
85 if plugin_base.exclude_tags or plugin_base.include_tags:
86 new_expr = " and ".join(
87 list(plugin_base.include_tags)
88 + [f"not {tag}" for tag in plugin_base.exclude_tags]
89 )
90
91 if config.option.markexpr:
92 config.option.markexpr += f" and {new_expr}"
93 else:
94 config.option.markexpr = new_expr
95
96 if config.pluginmanager.hasplugin("xdist"):
97 config.pluginmanager.register(XDistHooks())
98
99 if hasattr(config, "workerinput"):
100 plugin_base.restore_important_follower_config(config.workerinput)
101 plugin_base.configure_follower(config.workerinput["follower_ident"])
102 else:
103 if config.option.write_idents and os.path.exists(
104 config.option.write_idents
105 ):
106 os.remove(config.option.write_idents)
107
108 plugin_base.pre_begin(config.option)
109
110 plugin_base.set_coverage_flag(
111 bool(getattr(config.option, "cov_source", False))
112 )
113
114 plugin_base.set_fixture_functions(PytestFixtureFunctions)
115
116 if config.option.dump_pyannotate:
117 global DUMP_PYANNOTATE
118 DUMP_PYANNOTATE = True
119
120
121DUMP_PYANNOTATE = False
122
123
124@pytest.fixture(autouse=True)
125def collect_types_fixture():
126 if DUMP_PYANNOTATE:
127 from pyannotate_runtime import collect_types
128
129 collect_types.start()
130 yield
131 if DUMP_PYANNOTATE:
132 collect_types.stop()
133
134
135def _log_sqlalchemy_info(session):
136 import sqlalchemy
137 from sqlalchemy import __version__
138 from sqlalchemy.util import has_compiled_ext
139 from sqlalchemy.util._has_cy import _CYEXTENSION_MSG
140
141 greet = "sqlalchemy installation"
142 site = "no user site" if sys.flags.no_user_site else "user site loaded"
143 msgs = [
144 f"SQLAlchemy {__version__} ({site})",
145 f"Path: {sqlalchemy.__file__}",
146 ]
147
148 if has_compiled_ext():
149 from sqlalchemy.cyextension import util
150
151 msgs.append(f"compiled extension enabled, e.g. {util.__file__} ")
152 else:
153 msgs.append(f"compiled extension not enabled; {_CYEXTENSION_MSG}")
154
155 pm = session.config.pluginmanager.get_plugin("terminalreporter")
156 if pm:
157 pm.write_sep("=", greet)
158 for m in msgs:
159 pm.write_line(m)
160 else:
161 # fancy pants reporter not found, fallback to plain print
162 print("=" * 25, greet, "=" * 25)
163 for m in msgs:
164 print(m)
165
166
167def pytest_sessionstart(session):
168 from sqlalchemy.testing import asyncio
169
170 _log_sqlalchemy_info(session)
171 asyncio._assume_async(plugin_base.post_begin)
172
173
174def pytest_sessionfinish(session):
175 from sqlalchemy.testing import asyncio
176
177 asyncio._maybe_async_provisioning(plugin_base.final_process_cleanup)
178
179 if session.config.option.dump_pyannotate:
180 from pyannotate_runtime import collect_types
181
182 collect_types.dump_stats(session.config.option.dump_pyannotate)
183
184
185def pytest_unconfigure(config):
186 from sqlalchemy.testing import asyncio
187
188 asyncio._shutdown()
189
190
191def pytest_collection_finish(session):
192 if session.config.option.dump_pyannotate:
193 from pyannotate_runtime import collect_types
194
195 lib_sqlalchemy = os.path.abspath("lib/sqlalchemy")
196
197 def _filter(filename):
198 filename = os.path.normpath(os.path.abspath(filename))
199 if "lib/sqlalchemy" not in os.path.commonpath(
200 [filename, lib_sqlalchemy]
201 ):
202 return None
203 if "testing" in filename:
204 return None
205
206 return filename
207
208 collect_types.init_types_collection(filter_filename=_filter)
209
210
211class XDistHooks:
212 def pytest_configure_node(self, node):
213 from sqlalchemy.testing import provision
214 from sqlalchemy.testing import asyncio
215
216 # the master for each node fills workerinput dictionary
217 # which pytest-xdist will transfer to the subprocess
218
219 plugin_base.memoize_important_follower_config(node.workerinput)
220
221 node.workerinput["follower_ident"] = "test_%s" % uuid.uuid4().hex[0:12]
222
223 asyncio._maybe_async_provisioning(
224 provision.create_follower_db, node.workerinput["follower_ident"]
225 )
226
227 def pytest_testnodedown(self, node, error):
228 from sqlalchemy.testing import provision
229 from sqlalchemy.testing import asyncio
230
231 asyncio._maybe_async_provisioning(
232 provision.drop_follower_db, node.workerinput["follower_ident"]
233 )
234
235
236def pytest_collection_modifyitems(session, config, items):
237 # look for all those classes that specify __backend__ and
238 # expand them out into per-database test cases.
239
240 # this is much easier to do within pytest_pycollect_makeitem, however
241 # pytest is iterating through cls.__dict__ as makeitem is
242 # called which causes a "dictionary changed size" error on py3k.
243 # I'd submit a pullreq for them to turn it into a list first, but
244 # it's to suit the rather odd use case here which is that we are adding
245 # new classes to a module on the fly.
246
247 from sqlalchemy.testing import asyncio
248
249 rebuilt_items = collections.defaultdict(
250 lambda: collections.defaultdict(list)
251 )
252
253 items[:] = [
254 item
255 for item in items
256 if item.getparent(pytest.Class) is not None
257 and not item.getparent(pytest.Class).name.startswith("_")
258 ]
259
260 test_classes = {item.getparent(pytest.Class) for item in items}
261
262 def collect(element):
263 for inst_or_fn in element.collect():
264 if isinstance(inst_or_fn, pytest.Collector):
265 yield from collect(inst_or_fn)
266 else:
267 yield inst_or_fn
268
269 def setup_test_classes():
270 for test_class in test_classes:
271 # transfer legacy __backend__ and __sparse_backend__ symbols
272 # to be markers
273 if getattr(test_class.cls, "__backend__", False) or getattr(
274 test_class.cls, "__only_on__", False
275 ):
276 add_markers = {"backend"}
277 elif getattr(test_class.cls, "__sparse_backend__", False):
278 add_markers = {"sparse_backend", "backend"}
279 elif getattr(test_class.cls, "__sparse_driver_backend__", False):
280 add_markers = {"sparse_driver_backend", "backend"}
281 else:
282 add_markers = frozenset()
283
284 existing_markers = {
285 mark.name for mark in test_class.iter_markers()
286 }
287 add_markers = add_markers - existing_markers
288 all_markers = existing_markers.union(add_markers)
289
290 for marker in add_markers:
291 test_class.add_marker(marker)
292
293 sub_tests = list(
294 plugin_base.generate_sub_tests(
295 test_class.cls, test_class.module, all_markers
296 )
297 )
298 if not sub_tests:
299 rebuilt_items[test_class.cls]
300
301 for sub_cls in sub_tests:
302 if sub_cls is not test_class.cls:
303 per_cls_dict = rebuilt_items[test_class.cls]
304
305 module = test_class.getparent(pytest.Module)
306
307 new_cls = pytest.Class.from_parent(
308 name=sub_cls.__name__, parent=module
309 )
310 for marker in add_markers:
311 new_cls.add_marker(marker)
312
313 for fn in collect(new_cls):
314 per_cls_dict[fn.name].append(fn)
315
316 # class requirements will sometimes need to access the DB to check
317 # capabilities, so need to do this for async
318 asyncio._maybe_async_provisioning(setup_test_classes)
319
320 newitems = []
321 for item in items:
322 cls_ = item.cls
323 if cls_ in rebuilt_items:
324 newitems.extend(rebuilt_items[cls_][item.name])
325 else:
326 newitems.append(item)
327
328 # seems like the functions attached to a test class aren't sorted already?
329 # is that true and why's that? (when using unittest, they're sorted)
330 items[:] = sorted(
331 newitems,
332 key=lambda item: (
333 item.getparent(pytest.Module).name,
334 item.getparent(pytest.Class).name,
335 item.name,
336 ),
337 )
338
339
340def pytest_pycollect_makeitem(collector, name, obj):
341 if inspect.isclass(obj) and plugin_base.want_class(name, obj):
342 from sqlalchemy.testing import config
343
344 if config.any_async:
345 obj = _apply_maybe_async(obj)
346
347 return [
348 pytest.Class.from_parent(
349 name=parametrize_cls.__name__, parent=collector
350 )
351 for parametrize_cls in _parametrize_cls(collector.module, obj)
352 ]
353 elif (
354 inspect.isfunction(obj)
355 and collector.cls is not None
356 and plugin_base.want_method(collector.cls, obj)
357 ):
358 # None means, fall back to default logic, which includes
359 # method-level parametrize
360 return None
361 else:
362 # empty list means skip this item
363 return []
364
365
366def _is_wrapped_coroutine_function(fn):
367 while hasattr(fn, "__wrapped__"):
368 fn = fn.__wrapped__
369
370 return inspect.iscoroutinefunction(fn)
371
372
373def _apply_maybe_async(obj, recurse=True):
374 from sqlalchemy.testing import asyncio
375
376 for name, value in vars(obj).items():
377 if (
378 (callable(value) or isinstance(value, classmethod))
379 and not getattr(value, "_maybe_async_applied", False)
380 and (name.startswith("test_"))
381 and not _is_wrapped_coroutine_function(value)
382 ):
383 is_classmethod = False
384 if isinstance(value, classmethod):
385 value = value.__func__
386 is_classmethod = True
387
388 @_pytest_fn_decorator
389 def make_async(fn, *args, **kwargs):
390 return asyncio._maybe_async(fn, *args, **kwargs)
391
392 do_async = make_async(value)
393 if is_classmethod:
394 do_async = classmethod(do_async)
395 do_async._maybe_async_applied = True
396
397 setattr(obj, name, do_async)
398 if recurse:
399 for cls in obj.mro()[1:]:
400 if cls != object:
401 _apply_maybe_async(cls, False)
402 return obj
403
404
405def _parametrize_cls(module, cls):
406 """implement a class-based version of pytest parametrize."""
407
408 if "_sa_parametrize" not in cls.__dict__:
409 return [cls]
410
411 _sa_parametrize = cls._sa_parametrize
412 classes = []
413 for full_param_set in itertools.product(
414 *[params for argname, params in _sa_parametrize]
415 ):
416 cls_variables = {}
417
418 for argname, param in zip(
419 [_sa_param[0] for _sa_param in _sa_parametrize], full_param_set
420 ):
421 if not argname:
422 raise TypeError("need argnames for class-based combinations")
423 argname_split = re.split(r",\s*", argname)
424 for arg, val in zip(argname_split, param.values):
425 cls_variables[arg] = val
426 parametrized_name = "_".join(
427 re.sub(r"\W", "", token)
428 for param in full_param_set
429 for token in param.id.split("-")
430 )
431 name = "%s_%s" % (cls.__name__, parametrized_name)
432 newcls = type.__new__(type, name, (cls,), cls_variables)
433 setattr(module, name, newcls)
434 classes.append(newcls)
435 return classes
436
437
438_current_class = None
439
440_current_warning_context = None
441
442
443def pytest_runtest_setup(item):
444 from sqlalchemy.testing import asyncio
445
446 # pytest_runtest_setup runs *before* pytest fixtures with scope="class".
447 # plugin_base.start_test_class_outside_fixtures may opt to raise SkipTest
448 # for the whole class and has to run things that are across all current
449 # databases, so we run this outside of the pytest fixture system altogether
450 # and ensure asyncio greenlet if any engines are async
451
452 global _current_class, _current_warning_context
453
454 if isinstance(item, pytest.Function) and _current_class is None:
455 asyncio._maybe_async_provisioning(
456 plugin_base.start_test_class_outside_fixtures,
457 item.cls,
458 )
459 _current_class = item.getparent(pytest.Class)
460
461 if hasattr(_current_class.cls, "__warnings__"):
462 import warnings
463
464 _current_warning_context = warnings.catch_warnings()
465 _current_warning_context.__enter__()
466 for warning_message in _current_class.cls.__warnings__:
467 warnings.filterwarnings("ignore", warning_message)
468
469
470@pytest.hookimpl(hookwrapper=True)
471def pytest_runtest_teardown(item, nextitem):
472 # runs inside of pytest function fixture scope
473 # after test function runs
474
475 from sqlalchemy.testing import asyncio
476
477 asyncio._maybe_async(plugin_base.after_test, item)
478
479 yield
480 # this is now after all the fixture teardown have run, the class can be
481 # finalized. Since pytest v7 this finalizer can no longer be added in
482 # pytest_runtest_setup since the class has not yet been setup at that
483 # time.
484 # See https://github.com/pytest-dev/pytest/issues/9343
485
486 global _current_class, _current_report, _current_warning_context
487
488 if _current_class is not None and (
489 # last test or a new class
490 nextitem is None
491 or nextitem.getparent(pytest.Class) is not _current_class
492 ):
493
494 if _current_warning_context is not None:
495 _current_warning_context.__exit__(None, None, None)
496 _current_warning_context = None
497
498 _current_class = None
499
500 try:
501 asyncio._maybe_async_provisioning(
502 plugin_base.stop_test_class_outside_fixtures, item.cls
503 )
504 except Exception as e:
505 # in case of an exception during teardown attach the original
506 # error to the exception message, otherwise it will get lost
507 if _current_report.failed:
508 if not e.args:
509 e.args = (
510 "__Original test failure__:\n"
511 + _current_report.longreprtext,
512 )
513 elif e.args[-1] and isinstance(e.args[-1], str):
514 args = list(e.args)
515 args[-1] += (
516 "\n__Original test failure__:\n"
517 + _current_report.longreprtext
518 )
519 e.args = tuple(args)
520 else:
521 e.args += (
522 "__Original test failure__",
523 _current_report.longreprtext,
524 )
525 raise
526 finally:
527 _current_report = None
528
529
530def pytest_runtest_call(item):
531 # runs inside of pytest function fixture scope
532 # before test function runs
533
534 from sqlalchemy.testing import asyncio
535
536 asyncio._maybe_async(
537 plugin_base.before_test,
538 item,
539 item.module.__name__,
540 item.cls,
541 item.name,
542 )
543
544
545_current_report = None
546
547
548def pytest_runtest_logreport(report):
549 global _current_report
550 if report.when == "call":
551 _current_report = report
552
553
554@pytest.fixture(scope="class")
555def setup_class_methods(request):
556 from sqlalchemy.testing import asyncio
557
558 cls = request.cls
559
560 if hasattr(cls, "setup_test_class"):
561 asyncio._maybe_async(cls.setup_test_class)
562
563 yield
564
565 if hasattr(cls, "teardown_test_class"):
566 asyncio._maybe_async(cls.teardown_test_class)
567
568 asyncio._maybe_async(plugin_base.stop_test_class, cls)
569
570
571@pytest.fixture(scope="function")
572def setup_test_methods(request):
573 from sqlalchemy.testing import asyncio
574
575 # called for each test
576
577 self = request.instance
578
579 # before this fixture runs:
580
581 # 1. function level "autouse" fixtures under py3k (examples: TablesTest
582 # define tables / data, MappedTest define tables / mappers / data)
583
584 # 2. was for p2k. no longer applies
585
586 # 3. run outer xdist-style setup
587 if hasattr(self, "setup_test"):
588 asyncio._maybe_async(self.setup_test)
589
590 # alembic test suite is using setUp and tearDown
591 # xdist methods; support these in the test suite
592 # for the near term
593 if hasattr(self, "setUp"):
594 asyncio._maybe_async(self.setUp)
595
596 # inside the yield:
597 # 4. function level fixtures defined on test functions themselves,
598 # e.g. "connection", "metadata" run next
599
600 # 5. pytest hook pytest_runtest_call then runs
601
602 # 6. test itself runs
603
604 yield
605
606 # yield finishes:
607
608 # 7. function level fixtures defined on test functions
609 # themselves, e.g. "connection" rolls back the transaction, "metadata"
610 # emits drop all
611
612 # 8. pytest hook pytest_runtest_teardown hook runs, this is associated
613 # with fixtures close all sessions, provisioning.stop_test_class(),
614 # engines.testing_reaper -> ensure all connection pool connections
615 # are returned, engines created by testing_engine that aren't the
616 # config engine are disposed
617
618 asyncio._maybe_async(plugin_base.after_test_fixtures, self)
619
620 # 10. run xdist-style teardown
621 if hasattr(self, "tearDown"):
622 asyncio._maybe_async(self.tearDown)
623
624 if hasattr(self, "teardown_test"):
625 asyncio._maybe_async(self.teardown_test)
626
627 # 11. was for p2k. no longer applies
628
629 # 12. function level "autouse" fixtures under py3k (examples: TablesTest /
630 # MappedTest delete table data, possibly drop tables and clear mappers
631 # depending on the flags defined by the test class)
632
633
634def _pytest_fn_decorator(target):
635 """Port of langhelpers.decorator with pytest-specific tricks."""
636
637 from sqlalchemy.util.langhelpers import format_argspec_plus
638 from sqlalchemy.util.compat import inspect_getfullargspec
639
640 def _exec_code_in_env(code, env, fn_name):
641 # note this is affected by "from __future__ import annotations" at
642 # the top; exec'ed code will use non-evaluated annotations
643 # which allows us to be more flexible with code rendering
644 # in format_argpsec_plus()
645 exec(code, env)
646 return env[fn_name]
647
648 def decorate(fn, add_positional_parameters=()):
649 spec = inspect_getfullargspec(fn)
650 if add_positional_parameters:
651 spec.args.extend(add_positional_parameters)
652
653 metadata = dict(
654 __target_fn="__target_fn", __orig_fn="__orig_fn", name=fn.__name__
655 )
656 metadata.update(format_argspec_plus(spec, grouped=False))
657 code = (
658 """\
659def %(name)s%(grouped_args)s:
660 return %(__target_fn)s(%(__orig_fn)s, %(apply_kw)s)
661"""
662 % metadata
663 )
664 decorated = _exec_code_in_env(
665 code, {"__target_fn": target, "__orig_fn": fn}, fn.__name__
666 )
667 if not add_positional_parameters:
668 decorated.__defaults__ = getattr(fn, "__func__", fn).__defaults__
669 decorated.__wrapped__ = fn
670 return update_wrapper(decorated, fn)
671 else:
672 # this is the pytest hacky part. don't do a full update wrapper
673 # because pytest is really being sneaky about finding the args
674 # for the wrapped function
675 decorated.__module__ = fn.__module__
676 decorated.__name__ = fn.__name__
677 if hasattr(fn, "pytestmark"):
678 decorated.pytestmark = fn.pytestmark
679 return decorated
680
681 return decorate
682
683
684class PytestFixtureFunctions(plugin_base.FixtureFunctions):
685 def skip_test_exception(self, *arg, **kw):
686 return pytest.skip.Exception(*arg, **kw)
687
688 @property
689 def add_to_marker(self):
690 return pytest.mark
691
692 def mark_base_test_class(self):
693 return pytest.mark.usefixtures(
694 "setup_class_methods",
695 "setup_test_methods",
696 )
697
698 _combination_id_fns = {
699 "i": lambda obj: obj,
700 "r": repr,
701 "s": str,
702 "n": lambda obj: (
703 obj.__name__ if hasattr(obj, "__name__") else type(obj).__name__
704 ),
705 }
706
707 def combinations(self, *arg_sets, **kw):
708 """Facade for pytest.mark.parametrize.
709
710 Automatically derives argument names from the callable which in our
711 case is always a method on a class with positional arguments.
712
713 ids for parameter sets are derived using an optional template.
714
715 """
716 from sqlalchemy.testing import exclusions
717
718 if len(arg_sets) == 1 and hasattr(arg_sets[0], "__next__"):
719 arg_sets = list(arg_sets[0])
720
721 argnames = kw.pop("argnames", None)
722
723 def _filter_exclusions(args):
724 result = []
725 gathered_exclusions = []
726 for a in args:
727 if isinstance(a, exclusions.compound):
728 gathered_exclusions.append(a)
729 else:
730 result.append(a)
731
732 return result, gathered_exclusions
733
734 id_ = kw.pop("id_", None)
735
736 tobuild_pytest_params = []
737 has_exclusions = False
738 if id_:
739 _combination_id_fns = self._combination_id_fns
740
741 # because itemgetter is not consistent for one argument vs.
742 # multiple, make it multiple in all cases and use a slice
743 # to omit the first argument
744 _arg_getter = operator.itemgetter(
745 0,
746 *[
747 idx
748 for idx, char in enumerate(id_)
749 if char in ("n", "r", "s", "a")
750 ],
751 )
752 fns = [
753 (operator.itemgetter(idx), _combination_id_fns[char])
754 for idx, char in enumerate(id_)
755 if char in _combination_id_fns
756 ]
757
758 for arg in arg_sets:
759 if not isinstance(arg, tuple):
760 arg = (arg,)
761
762 fn_params, param_exclusions = _filter_exclusions(arg)
763
764 parameters = _arg_getter(fn_params)[1:]
765
766 if param_exclusions:
767 has_exclusions = True
768
769 tobuild_pytest_params.append(
770 (
771 parameters,
772 param_exclusions,
773 "-".join(
774 comb_fn(getter(arg)) for getter, comb_fn in fns
775 ),
776 )
777 )
778
779 else:
780 for arg in arg_sets:
781 if not isinstance(arg, tuple):
782 arg = (arg,)
783
784 fn_params, param_exclusions = _filter_exclusions(arg)
785
786 if param_exclusions:
787 has_exclusions = True
788
789 tobuild_pytest_params.append(
790 (fn_params, param_exclusions, None)
791 )
792
793 pytest_params = []
794 for parameters, param_exclusions, id_ in tobuild_pytest_params:
795 if has_exclusions:
796 parameters += (param_exclusions,)
797
798 param = pytest.param(*parameters, id=id_)
799 pytest_params.append(param)
800
801 def decorate(fn):
802 if inspect.isclass(fn):
803 if has_exclusions:
804 raise NotImplementedError(
805 "exclusions not supported for class level combinations"
806 )
807 if "_sa_parametrize" not in fn.__dict__:
808 fn._sa_parametrize = []
809 fn._sa_parametrize.append((argnames, pytest_params))
810 return fn
811 else:
812 _fn_argnames = inspect.getfullargspec(fn).args[1:]
813 if argnames is None:
814 _argnames = _fn_argnames
815 else:
816 _argnames = re.split(r", *", argnames)
817
818 if has_exclusions:
819 existing_exl = sum(
820 1 for n in _fn_argnames if n.startswith("_exclusions")
821 )
822 current_exclusion_name = f"_exclusions_{existing_exl}"
823 _argnames += [current_exclusion_name]
824
825 @_pytest_fn_decorator
826 def check_exclusions(fn, *args, **kw):
827 _exclusions = args[-1]
828 if _exclusions:
829 exlu = exclusions.compound().add(*_exclusions)
830 fn = exlu(fn)
831 return fn(*args[:-1], **kw)
832
833 fn = check_exclusions(
834 fn, add_positional_parameters=(current_exclusion_name,)
835 )
836
837 return pytest.mark.parametrize(_argnames, pytest_params)(fn)
838
839 return decorate
840
841 def param_ident(self, *parameters):
842 ident = parameters[0]
843 return pytest.param(*parameters[1:], id=ident)
844
845 def fixture(self, *arg, **kw):
846 from sqlalchemy.testing import config
847 from sqlalchemy.testing import asyncio
848
849 # wrapping pytest.fixture function. determine if
850 # decorator was called as @fixture or @fixture().
851 if len(arg) > 0 and callable(arg[0]):
852 # was called as @fixture(), we have the function to wrap.
853 fn = arg[0]
854 arg = arg[1:]
855 else:
856 # was called as @fixture, don't have the function yet.
857 fn = None
858
859 # create a pytest.fixture marker. because the fn is not being
860 # passed, this is always a pytest.FixtureFunctionMarker()
861 # object (or whatever pytest is calling it when you read this)
862 # that is waiting for a function.
863 fixture = pytest.fixture(*arg, **kw)
864
865 # now apply wrappers to the function, including fixture itself
866
867 def wrap(fn):
868 if config.any_async:
869 fn = asyncio._maybe_async_wrapper(fn)
870 # other wrappers may be added here
871
872 # now apply FixtureFunctionMarker
873 fn = fixture(fn)
874
875 return fn
876
877 if fn:
878 return wrap(fn)
879 else:
880 return wrap
881
882 def get_current_test_name(self):
883 return os.environ.get("PYTEST_CURRENT_TEST")
884
885 def async_test(self, fn):
886 from sqlalchemy.testing import asyncio
887
888 @_pytest_fn_decorator
889 def decorate(fn, *args, **kwargs):
890 asyncio._run_coroutine_function(fn, *args, **kwargs)
891
892 return decorate(fn)
893 