Team Ai
Datasetpublic

codekingpro/portable-devtools

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