codekingpro/portable-devtools
115k
1# testing/util.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
12from collections import deque
13from collections import namedtuple
14import contextlib
15import decimal
16import gc
17from itertools import chain
18import pickle
19import random
20import sys
21from sys import getsizeof
22import time
23import types
24from typing import Any
25
26from . import config
27from . import mock
28from .. import inspect
29from ..engine import Connection
30from ..schema import Column
31from ..schema import DropConstraint
32from ..schema import DropTable
33from ..schema import ForeignKeyConstraint
34from ..schema import MetaData
35from ..schema import Table
36from ..sql import schema
37from ..sql.sqltypes import Integer
38from ..util import decorator
39from ..util import defaultdict
40from ..util import has_refcount_gc
41from ..util import inspect_getfullargspec
42
43
44if not has_refcount_gc:
45
46 def non_refcount_gc_collect(*args):
47 gc.collect()
48 gc.collect()
49
50 gc_collect = lazy_gc = non_refcount_gc_collect
51else:
52 # assume CPython - straight gc.collect, lazy_gc() is a pass
53 gc_collect = gc.collect
54
55 def lazy_gc():
56 pass
57
58
59def picklers():
60 nt = namedtuple("picklers", ["loads", "dumps"])
61
62 for protocol in range(-2, pickle.HIGHEST_PROTOCOL + 1):
63 yield nt(pickle.loads, lambda d: pickle.dumps(d, protocol))
64
65
66def random_choices(population, k=1):
67 return random.choices(population, k=k)
68
69
70def round_decimal(value, prec):
71 if isinstance(value, float):
72 return round(value, prec)
73
74 # can also use shift() here but that is 2.6 only
75 return (value * decimal.Decimal("1" + "0" * prec)).to_integral(
76 decimal.ROUND_FLOOR
77 ) / pow(10, prec)
78
79
80class RandomSet(set):
81 def __iter__(self):
82 l = list(set.__iter__(self))
83 random.shuffle(l)
84 return iter(l)
85
86 def pop(self):
87 index = random.randint(0, len(self) - 1)
88 item = list(set.__iter__(self))[index]
89 self.remove(item)
90 return item
91
92 def union(self, other):
93 return RandomSet(set.union(self, other))
94
95 def difference(self, other):
96 return RandomSet(set.difference(self, other))
97
98 def intersection(self, other):
99 return RandomSet(set.intersection(self, other))
100
101 def copy(self):
102 return RandomSet(self)
103
104
105def conforms_partial_ordering(tuples, sorted_elements):
106 """True if the given sorting conforms to the given partial ordering."""
107
108 deps = defaultdict(set)
109 for parent, child in tuples:
110 deps[parent].add(child)
111 for i, node in enumerate(sorted_elements):
112 for n in sorted_elements[i:]:
113 if node in deps[n]:
114 return False
115 else:
116 return True
117
118
119def all_partial_orderings(tuples, elements):
120 edges = defaultdict(set)
121 for parent, child in tuples:
122 edges[child].add(parent)
123
124 def _all_orderings(elements):
125 if len(elements) == 1:
126 yield list(elements)
127 else:
128 for elem in elements:
129 subset = set(elements).difference([elem])
130 if not subset.intersection(edges[elem]):
131 for sub_ordering in _all_orderings(subset):
132 yield [elem] + sub_ordering
133
134 return iter(_all_orderings(elements))
135
136
137def function_named(fn, name):
138 """Return a function with a given __name__.
139
140 Will assign to __name__ and return the original function if possible on
141 the Python implementation, otherwise a new function will be constructed.
142
143 This function should be phased out as much as possible
144 in favor of @decorator. Tests that "generate" many named tests
145 should be modernized.
146
147 """
148 try:
149 fn.__name__ = name
150 except TypeError:
151 fn = types.FunctionType(
152 fn.__code__, fn.__globals__, name, fn.__defaults__, fn.__closure__
153 )
154 return fn
155
156
157def run_as_contextmanager(ctx, fn, *arg, **kw):
158 """Run the given function under the given contextmanager,
159 simulating the behavior of 'with' to support older
160 Python versions.
161
162 This is not necessary anymore as we have placed 2.6
163 as minimum Python version, however some tests are still using
164 this structure.
165
166 """
167
168 obj = ctx.__enter__()
169 try:
170 result = fn(obj, *arg, **kw)
171 ctx.__exit__(None, None, None)
172 return result
173 except:
174 exc_info = sys.exc_info()
175 raise_ = ctx.__exit__(*exc_info)
176 if not raise_:
177 raise
178 else:
179 return raise_
180
181
182def rowset(results):
183 """Converts the results of sql execution into a plain set of column tuples.
184
185 Useful for asserting the results of an unordered query.
186 """
187
188 return {tuple(row) for row in results}
189
190
191def fail(msg):
192 assert False, msg
193
194
195@decorator
196def provide_metadata(fn, *args, **kw):
197 """Provide bound MetaData for a single test, dropping afterwards.
198
199 Legacy; use the "metadata" pytest fixture.
200
201 """
202
203 from . import fixtures
204
205 metadata = schema.MetaData()
206 self = args[0]
207 prev_meta = getattr(self, "metadata", None)
208 self.metadata = metadata
209 try:
210 return fn(*args, **kw)
211 finally:
212 # close out some things that get in the way of dropping tables.
213 # when using the "metadata" fixture, there is a set ordering
214 # of things that makes sure things are cleaned up in order, however
215 # the simple "decorator" nature of this legacy function means
216 # we have to hardcode some of that cleanup ahead of time.
217
218 # close ORM sessions
219 fixtures.close_all_sessions()
220
221 # integrate with the "connection" fixture as there are many
222 # tests where it is used along with provide_metadata
223 cfc = fixtures.base._connection_fixture_connection
224 if cfc:
225 # TODO: this warning can be used to find all the places
226 # this is used with connection fixture
227 # warn("mixing legacy provide metadata with connection fixture")
228 drop_all_tables_from_metadata(metadata, cfc)
229 # as the provide_metadata fixture is often used with "testing.db",
230 # when we do the drop we have to commit the transaction so that
231 # the DB is actually updated as the CREATE would have been
232 # committed
233 cfc.get_transaction().commit()
234 else:
235 drop_all_tables_from_metadata(metadata, config.db)
236 self.metadata = prev_meta
237
238
239def flag_combinations(*combinations):
240 """A facade around @testing.combinations() oriented towards boolean
241 keyword-based arguments.
242
243 Basically generates a nice looking identifier based on the keywords
244 and also sets up the argument names.
245
246 E.g.::
247
248 @testing.flag_combinations(
249 dict(lazy=False, passive=False),
250 dict(lazy=True, passive=False),
251 dict(lazy=False, passive=True),
252 dict(lazy=False, passive=True, raiseload=True),
253 )
254 def test_fn(lazy, passive, raiseload): ...
255
256 would result in::
257
258 @testing.combinations(
259 ("", False, False, False),
260 ("lazy", True, False, False),
261 ("lazy_passive", True, True, False),
262 ("lazy_passive", True, True, True),
263 id_="iaaa",
264 argnames="lazy,passive,raiseload",
265 )
266 def test_fn(lazy, passive, raiseload): ...
267
268 """
269
270 keys = set()
271
272 for d in combinations:
273 keys.update(d)
274
275 keys = sorted(keys)
276
277 return config.combinations(
278 *[
279 ("_".join(k for k in keys if d.get(k, False)),)
280 + tuple(d.get(k, False) for k in keys)
281 for d in combinations
282 ],
283 id_="i" + ("a" * len(keys)),
284 argnames=",".join(keys),
285 )
286
287
288def lambda_combinations(lambda_arg_sets, **kw):
289 args = inspect_getfullargspec(lambda_arg_sets)
290
291 arg_sets = lambda_arg_sets(*[mock.Mock() for arg in args[0]])
292
293 def create_fixture(pos):
294 def fixture(**kw):
295 return lambda_arg_sets(**kw)[pos]
296
297 fixture.__name__ = "fixture_%3.3d" % pos
298 return fixture
299
300 return config.combinations(
301 *[(create_fixture(i),) for i in range(len(arg_sets))], **kw
302 )
303
304
305def resolve_lambda(__fn, **kw):
306 """Given a no-arg lambda and a namespace, return a new lambda that
307 has all the values filled in.
308
309 This is used so that we can have module-level fixtures that
310 refer to instance-level variables using lambdas.
311
312 """
313
314 pos_args = inspect_getfullargspec(__fn)[0]
315 pass_pos_args = {arg: kw.pop(arg) for arg in pos_args}
316 glb = dict(__fn.__globals__)
317 glb.update(kw)
318 new_fn = types.FunctionType(__fn.__code__, glb)
319 return new_fn(**pass_pos_args)
320
321
322def metadata_fixture(ddl="function"):
323 """Provide MetaData for a pytest fixture."""
324
325 def decorate(fn):
326 def run_ddl(self):
327 metadata = self.metadata = schema.MetaData()
328 try:
329 result = fn(self, metadata)
330 metadata.create_all(config.db)
331 # TODO:
332 # somehow get a per-function dml erase fixture here
333 yield result
334 finally:
335 metadata.drop_all(config.db)
336
337 return config.fixture(scope=ddl)(run_ddl)
338
339 return decorate
340
341
342def force_drop_names(*names):
343 """Force the given table names to be dropped after test complete,
344 isolating for foreign key cycles
345
346 """
347
348 @decorator
349 def go(fn, *args, **kw):
350 try:
351 return fn(*args, **kw)
352 finally:
353 drop_all_tables(config.db, inspect(config.db), include_names=names)
354
355 return go
356
357
358class adict(dict):
359 """Dict keys available as attributes. Shadows."""
360
361 def __getattribute__(self, key):
362 try:
363 return self[key]
364 except KeyError:
365 return dict.__getattribute__(self, key)
366
367 def __call__(self, *keys):
368 return tuple([self[key] for key in keys])
369
370 get_all = __call__
371
372
373def drop_all_tables_from_metadata(metadata, engine_or_connection):
374 from . import engines
375
376 def go(connection):
377 engines.testing_reaper.prepare_for_drop_tables(connection)
378
379 if not connection.dialect.supports_alter:
380 from . import assertions
381
382 with assertions.expect_warnings(
383 "Can't sort tables", assert_=False
384 ):
385 metadata.drop_all(connection)
386 else:
387 metadata.drop_all(connection)
388
389 if not isinstance(engine_or_connection, Connection):
390 with engine_or_connection.begin() as connection:
391 go(connection)
392 else:
393 go(engine_or_connection)
394
395
396def drop_all_tables(
397 engine,
398 inspector,
399 schema=None,
400 consider_schemas=(None,),
401 include_names=None,
402):
403 if include_names is not None:
404 include_names = set(include_names)
405
406 if schema is not None:
407 assert consider_schemas == (
408 None,
409 ), "consider_schemas and schema are mutually exclusive"
410 consider_schemas = (schema,)
411
412 with engine.begin() as conn:
413 for table_key, fkcs in reversed(
414 inspector.sort_tables_on_foreign_key_dependency(
415 consider_schemas=consider_schemas
416 )
417 ):
418 if table_key:
419 if (
420 include_names is not None
421 and table_key[1] not in include_names
422 ):
423 continue
424 conn.execute(
425 DropTable(
426 Table(table_key[1], MetaData(), schema=table_key[0])
427 )
428 )
429 elif fkcs:
430 if not engine.dialect.supports_alter:
431 continue
432 for t_key, fkc in fkcs:
433 if (
434 include_names is not None
435 and t_key[1] not in include_names
436 ):
437 continue
438 tb = Table(
439 t_key[1],
440 MetaData(),
441 Column("x", Integer),
442 Column("y", Integer),
443 schema=t_key[0],
444 )
445 conn.execute(
446 DropConstraint(
447 ForeignKeyConstraint([tb.c.x], [tb.c.y], name=fkc)
448 )
449 )
450
451
452def teardown_events(event_cls):
453 @decorator
454 def decorate(fn, *arg, **kw):
455 try:
456 return fn(*arg, **kw)
457 finally:
458 event_cls._clear()
459
460 return decorate
461
462
463def total_size(o):
464 """Returns the approximate memory footprint an object and all of its
465 contents.
466
467 source: https://code.activestate.com/recipes/577504/
468
469
470 """
471
472 def dict_handler(d):
473 return chain.from_iterable(d.items())
474
475 all_handlers = {
476 tuple: iter,
477 list: iter,
478 deque: iter,
479 dict: dict_handler,
480 set: iter,
481 frozenset: iter,
482 }
483 seen = set() # track which object id's have already been seen
484 default_size = getsizeof(0) # estimate sizeof object without __sizeof__
485
486 def sizeof(o):
487 if id(o) in seen: # do not double count the same object
488 return 0
489 seen.add(id(o))
490 s = getsizeof(o, default_size)
491
492 for typ, handler in all_handlers.items():
493 if isinstance(o, typ):
494 s += sum(map(sizeof, handler(o)))
495 break
496 return s
497
498 return sizeof(o)
499
500
501def count_cache_key_tuples(tup):
502 """given a cache key tuple, counts how many instances of actual
503 tuples are found.
504
505 used to alert large jumps in cache key complexity.
506
507 """
508 stack = [tup]
509
510 sentinel = object()
511 num_elements = 0
512
513 while stack:
514 elem = stack.pop(0)
515 if elem is sentinel:
516 num_elements += 1
517 elif isinstance(elem, tuple):
518 if elem:
519 stack = list(elem) + [sentinel] + stack
520 return num_elements
521
522
523@contextlib.contextmanager
524def skip_if_timeout(seconds: float, cleanup: Any = None):
525
526 now = time.time()
527 yield
528 sec = time.time() - now
529 if sec > seconds:
530 try:
531 cleanup()
532 finally:
533 config.skip_test(
534 f"test took too long ({sec:.4f} seconds > {seconds})"
535 )
536 