codekingpro/portable-devtools
115k
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 