Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
exclusions.py477 linesDownload Raw Back to testing
1# testing/exclusions.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
9import contextlib
10import operator
11import re
12import sys
13
14from . import config
15from .. import util
16from ..util import decorator
17from ..util.compat import inspect_getfullargspec
18
19
20def skip_if(predicate, reason=None):
21    rule = compound()
22    pred = _as_predicate(predicate, reason)
23    rule.skips.add(pred)
24    return rule
25
26
27def fails_if(predicate, reason=None):
28    rule = compound()
29    pred = _as_predicate(predicate, reason)
30    rule.fails.add(pred)
31    return rule
32
33
34def warns_if(predicate, expression, assert_):
35    rule = compound()
36    pred = _as_predicate(predicate)
37    rule.warns[pred] = (expression, assert_)
38    return rule
39
40
41class compound:
42    def __init__(self):
43        self.fails = set()
44        self.skips = set()
45        self.warns = {}
46
47    def __add__(self, other):
48        return self.add(other)
49
50    def as_skips(self):
51        rule = compound()
52        rule.skips.update(self.skips)
53        rule.skips.update(self.fails)
54        return rule
55
56    def add(self, *others):
57        copy = compound()
58        copy.fails.update(self.fails)
59        copy.skips.update(self.skips)
60        copy.warns.update(self.warns)
61
62        for other in others:
63            copy.fails.update(other.fails)
64            copy.skips.update(other.skips)
65            copy.warns.update(other.warns)
66        return copy
67
68    def not_(self):
69        copy = compound()
70        copy.fails.update(NotPredicate(fail) for fail in self.fails)
71        copy.skips.update(NotPredicate(skip) for skip in self.skips)
72        copy.warns.update(
73            {
74                NotPredicate(warn): element
75                for warn, element in self.warns.items()
76            }
77        )
78        return copy
79
80    @property
81    def enabled(self):
82        return self.enabled_for_config(config._current)
83
84    def enabled_for_config(self, config):
85        for predicate in self.skips.union(self.fails):
86            if predicate(config):
87                return False
88        else:
89            return True
90
91    def matching_warnings(self, config):
92        return [
93            message
94            for predicate, (message, assert_) in self.warns.items()
95            if predicate(config)
96        ]
97
98    def matching_config_reasons(self, config):
99        return [
100            predicate._as_string(config)
101            for predicate in self.skips.union(self.fails)
102            if predicate(config)
103        ]
104
105    def _extend(self, other):
106        self.skips.update(other.skips)
107        self.fails.update(other.fails)
108        self.warns.update(other.warns)
109
110    def __call__(self, fn):
111        if hasattr(fn, "_sa_exclusion_extend"):
112            fn._sa_exclusion_extend._extend(self)
113            return fn
114
115        @decorator
116        def decorate(fn, *args, **kw):
117            return self._do(config._current, fn, *args, **kw)
118
119        decorated = decorate(fn)
120        decorated._sa_exclusion_extend = self
121        return decorated
122
123    @contextlib.contextmanager
124    def fail_if(self):
125        all_fails = compound()
126        all_fails.fails.update(self.skips.union(self.fails))
127
128        try:
129            yield
130        except Exception as ex:
131            all_fails._expect_failure(config._current, ex)
132        else:
133            all_fails._expect_success(config._current)
134
135    def _do(self, cfg, fn, *args, **kw):
136        for skip in self.skips:
137            if skip(cfg):
138                msg = "'%s' : %s" % (
139                    config.get_current_test_name(),
140                    skip._as_string(cfg),
141                )
142                config.skip_test(msg)
143
144        if self.warns:
145            from .assertions import expect_warnings
146
147            @contextlib.contextmanager
148            def _expect_warnings():
149                with contextlib.ExitStack() as stack:
150                    for expression, assert_ in self.warns.values():
151                        stack.enter_context(
152                            expect_warnings(expression, assert_=assert_)
153                        )
154                    yield
155
156            ctx = _expect_warnings()
157        else:
158            ctx = contextlib.nullcontext()
159
160        try:
161            with ctx:
162                return_value = fn(*args, **kw)
163        except Exception as ex:
164            self._expect_failure(cfg, ex, name=fn.__name__)
165        else:
166            self._expect_success(cfg, name=fn.__name__)
167            return return_value
168
169    def _expect_failure(self, config, ex, name="block"):
170        for fail in self.fails:
171            if fail(config):
172                print(
173                    "%s failed as expected (%s): %s "
174                    % (name, fail._as_string(config), ex)
175                )
176                break
177        else:
178            raise ex.with_traceback(sys.exc_info()[2])
179
180    def _expect_success(self, config, name="block"):
181        if not self.fails:
182            return
183
184        for fail in self.fails:
185            if fail(config):
186                raise AssertionError(
187                    "Unexpected success for '%s' (%s)"
188                    % (
189                        name,
190                        " and ".join(
191                            fail._as_string(config) for fail in self.fails
192                        ),
193                    )
194                )
195
196
197def only_if(predicate, reason=None):
198    predicate = _as_predicate(predicate)
199    return skip_if(NotPredicate(predicate), reason)
200
201
202def succeeds_if(predicate, reason=None):
203    predicate = _as_predicate(predicate)
204    return fails_if(NotPredicate(predicate), reason)
205
206
207class Predicate:
208    @classmethod
209    def as_predicate(cls, predicate, description=None):
210        if isinstance(predicate, compound):
211            return cls.as_predicate(predicate.enabled_for_config, description)
212        elif isinstance(predicate, Predicate):
213            if description and predicate.description is None:
214                predicate.description = description
215            return predicate
216        elif isinstance(predicate, (list, set)):
217            return OrPredicate(
218                [cls.as_predicate(pred) for pred in predicate], description
219            )
220        elif isinstance(predicate, tuple):
221            return SpecPredicate(*predicate)
222        elif isinstance(predicate, str):
223            tokens = re.match(
224                r"([\+\w]+)\s*(?:(>=|==|!=|<=|<|>)\s*([\d\.]+))?", predicate
225            )
226            if not tokens:
227                raise ValueError(
228                    "Couldn't locate DB name in predicate: %r" % predicate
229                )
230            db = tokens.group(1)
231            op = tokens.group(2)
232            spec = (
233                tuple(int(d) for d in tokens.group(3).split("."))
234                if tokens.group(3)
235                else None
236            )
237
238            return SpecPredicate(db, op, spec, description=description)
239        elif callable(predicate):
240            return LambdaPredicate(predicate, description)
241        else:
242            assert False, "unknown predicate type: %s" % predicate
243
244    def _format_description(self, config, negate=False):
245        bool_ = self(config)
246        if negate:
247            bool_ = not negate
248        return self.description % {
249            "driver": (
250                config.db.url.get_driver_name() if config else "<no driver>"
251            ),
252            "database": (
253                config.db.url.get_backend_name() if config else "<no database>"
254            ),
255            "doesnt_support": "doesn't support" if bool_ else "does support",
256            "does_support": "does support" if bool_ else "doesn't support",
257        }
258
259    def _as_string(self, config=None, negate=False):
260        raise NotImplementedError()
261
262
263class BooleanPredicate(Predicate):
264    def __init__(self, value, description=None):
265        self.value = value
266        self.description = description or "boolean %s" % value
267
268    def __call__(self, config):
269        return self.value
270
271    def _as_string(self, config, negate=False):
272        return self._format_description(config, negate=negate)
273
274
275class SpecPredicate(Predicate):
276    def __init__(self, db, op=None, spec=None, description=None):
277        self.db = db
278        self.op = op
279        self.spec = spec
280        self.description = description
281
282    _ops = {
283        "<": operator.lt,
284        ">": operator.gt,
285        "==": operator.eq,
286        "!=": operator.ne,
287        "<=": operator.le,
288        ">=": operator.ge,
289        "in": operator.contains,
290        "between": lambda val, pair: val >= pair[0] and val <= pair[1],
291    }
292
293    def __call__(self, config):
294        if config is None:
295            return False
296
297        engine = config.db
298
299        if "+" in self.db:
300            dialect, driver = self.db.split("+")
301        else:
302            dialect, driver = self.db, None
303
304        if dialect and engine.name != dialect:
305            return False
306        if driver is not None and engine.driver != driver:
307            return False
308
309        if self.op is not None:
310            assert driver is None, "DBAPI version specs not supported yet"
311
312            version = _server_version(engine)
313            oper = (
314                hasattr(self.op, "__call__") and self.op or self._ops[self.op]
315            )
316            return oper(version, self.spec)
317        else:
318            return True
319
320    def _as_string(self, config, negate=False):
321        if self.description is not None:
322            return self._format_description(config)
323        elif self.op is None:
324            if negate:
325                return "not %s" % self.db
326            else:
327                return "%s" % self.db
328        else:
329            if negate:
330                return "not %s %s %s" % (self.db, self.op, self.spec)
331            else:
332                return "%s %s %s" % (self.db, self.op, self.spec)
333
334
335class LambdaPredicate(Predicate):
336    def __init__(self, lambda_, description=None, args=None, kw=None):
337        spec = inspect_getfullargspec(lambda_)
338        if not spec[0]:
339            self.lambda_ = lambda db: lambda_()
340        else:
341            self.lambda_ = lambda_
342        self.args = args or ()
343        self.kw = kw or {}
344        if description:
345            self.description = description
346        elif lambda_.__doc__:
347            self.description = lambda_.__doc__
348        else:
349            self.description = "custom function"
350
351    def __call__(self, config):
352        return self.lambda_(config)
353
354    def _as_string(self, config, negate=False):
355        return self._format_description(config)
356
357
358class NotPredicate(Predicate):
359    def __init__(self, predicate, description=None):
360        self.predicate = predicate
361        self.description = description
362
363    def __call__(self, config):
364        return not self.predicate(config)
365
366    def _as_string(self, config, negate=False):
367        if self.description:
368            return self._format_description(config, not negate)
369        else:
370            return self.predicate._as_string(config, not negate)
371
372
373class OrPredicate(Predicate):
374    def __init__(self, predicates, description=None):
375        self.predicates = predicates
376        self.description = description
377
378    def __call__(self, config):
379        for pred in self.predicates:
380            if pred(config):
381                return True
382        return False
383
384    def _eval_str(self, config, negate=False):
385        if negate:
386            conjunction = " and "
387        else:
388            conjunction = " or "
389        return conjunction.join(
390            p._as_string(config, negate=negate) for p in self.predicates
391        )
392
393    def _negation_str(self, config):
394        if self.description is not None:
395            return "Not " + self._format_description(config)
396        else:
397            return self._eval_str(config, negate=True)
398
399    def _as_string(self, config, negate=False):
400        if negate:
401            return self._negation_str(config)
402        else:
403            if self.description is not None:
404                return self._format_description(config)
405            else:
406                return self._eval_str(config)
407
408
409_as_predicate = Predicate.as_predicate
410
411
412def _is_excluded(db, op, spec):
413    return SpecPredicate(db, op, spec)(config._current)
414
415
416def _server_version(engine):
417    """Return a server_version_info tuple."""
418
419    # force metadata to be retrieved
420    conn = engine.connect()
421    version = getattr(engine.dialect, "server_version_info", None)
422    if version is None:
423        version = ()
424    conn.close()
425    return version
426
427
428def db_spec(*dbs):
429    return OrPredicate([Predicate.as_predicate(db) for db in dbs])
430
431
432def open():  # noqa
433    return skip_if(BooleanPredicate(False, "mark as execute"))
434
435
436def closed(reason="marked as skip"):
437    return skip_if(BooleanPredicate(True, reason))
438
439
440def fails(reason=None):
441    return fails_if(BooleanPredicate(True, reason or "expected to fail"))
442
443
444def future():
445    return fails_if(BooleanPredicate(True, "Future feature"))
446
447
448def fails_on(db, reason=None):
449    return fails_if(db, reason)
450
451
452def fails_on_everything_except(*dbs):
453    return succeeds_if(OrPredicate([Predicate.as_predicate(db) for db in dbs]))
454
455
456def skip(db, reason=None):
457    return skip_if(db, reason)
458
459
460def only_on(dbs, reason=None):
461    return only_if(
462        OrPredicate(
463            [Predicate.as_predicate(db, reason) for db in util.to_list(dbs)]
464        )
465    )
466
467
468def exclude(db, op, spec, reason=None):
469    return skip_if(SpecPredicate(db, op, spec), reason)
470
471
472def against(config, *queries):
473    assert queries, "no queries sent!"
474    return OrPredicate([Predicate.as_predicate(query) for query in queries])(
475        config
476    )
477 
codekingpro/portable-devtools · Team Ai