codekingpro/portable-devtools
115k
1# testing/util.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
9
10from __future__ import annotations
11
12from collections import deque
13import decimal
14import gc
15from itertools import chain
16import random
17import sys
18from sys import getsizeof
19import types
20
21from . import config
22from . import mock
23from .. import inspect
24from ..engine import Connection
25from ..schema import Column
26from ..schema import DropConstraint
27from ..schema import DropTable
28from ..schema import ForeignKeyConstraint
29from ..schema import MetaData
30from ..schema import Table
31from ..sql import schema
32from ..sql.sqltypes import Integer
33from ..util import decorator
34from ..util import defaultdict
35from ..util import has_refcount_gc
36from ..util import inspect_getfullargspec
37
38
39if not has_refcount_gc:
40
41 def non_refcount_gc_collect(*args):
42 gc.collect()
43 gc.collect()
44
45 gc_collect = lazy_gc = non_refcount_gc_collect
46else:
47 # assume CPython - straight gc.collect, lazy_gc() is a pass
48 gc_collect = gc.collect
49
50 def lazy_gc():
51 pass
52
53
54def picklers():
55 picklers = set()
56 import pickle
57
58 picklers.add(pickle)
59
60 # yes, this thing needs this much testing
61 for pickle_ in picklers:
62 for protocol in range(-2, pickle.HIGHEST_PROTOCOL + 1):
63 yield 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
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
267 """
268
269 keys = set()
270
271 for d in combinations:
272 keys.update(d)
273
274 keys = sorted(keys)
275
276 return config.combinations(
277 *[
278 ("_".join(k for k in keys if d.get(k, False)),)
279 + tuple(d.get(k, False) for k in keys)
280 for d in combinations
281 ],
282 id_="i" + ("a" * len(keys)),
283 argnames=",".join(keys),
284 )
285
286
287def lambda_combinations(lambda_arg_sets, **kw):
288 args = inspect_getfullargspec(lambda_arg_sets)
289
290 arg_sets = lambda_arg_sets(*[mock.Mock() for arg in args[0]])
291
292 def create_fixture(pos):
293 def fixture(**kw):
294 return lambda_arg_sets(**kw)[pos]
295
296 fixture.__name__ = "fixture_%3.3d" % pos
297 return fixture
298
299 return config.combinations(
300 *[(create_fixture(i),) for i in range(len(arg_sets))], **kw
301 )
302
303
304def resolve_lambda(__fn, **kw):
305 """Given a no-arg lambda and a namespace, return a new lambda that
306 has all the values filled in.
307
308 This is used so that we can have module-level fixtures that
309 refer to instance-level variables using lambdas.
310
311 """
312
313 pos_args = inspect_getfullargspec(__fn)[0]
314 pass_pos_args = {arg: kw.pop(arg) for arg in pos_args}
315 glb = dict(__fn.__globals__)
316 glb.update(kw)
317 new_fn = types.FunctionType(__fn.__code__, glb)
318 return new_fn(**pass_pos_args)
319
320
321def metadata_fixture(ddl="function"):
322 """Provide MetaData for a pytest fixture."""
323
324 def decorate(fn):
325 def run_ddl(self):
326 metadata = self.metadata = schema.MetaData()
327 try:
328 result = fn(self, metadata)
329 metadata.create_all(config.db)
330 # TODO:
331 # somehow get a per-function dml erase fixture here
332 yield result
333 finally:
334 metadata.drop_all(config.db)
335
336 return config.fixture(scope=ddl)(run_ddl)
337
338 return decorate
339
340
341def force_drop_names(*names):
342 """Force the given table names to be dropped after test complete,
343 isolating for foreign key cycles
344
345 """
346
347 @decorator
348 def go(fn, *args, **kw):
349 try:
350 return fn(*args, **kw)
351 finally:
352 drop_all_tables(config.db, inspect(config.db), include_names=names)
353
354 return go
355
356
357class adict(dict):
358 """Dict keys available as attributes. Shadows."""
359
360 def __getattribute__(self, key):
361 try:
362 return self[key]
363 except KeyError:
364 return dict.__getattribute__(self, key)
365
366 def __call__(self, *keys):
367 return tuple([self[key] for key in keys])
368
369 get_all = __call__
370
371
372def drop_all_tables_from_metadata(metadata, engine_or_connection):
373 from . import engines
374
375 def go(connection):
376 engines.testing_reaper.prepare_for_drop_tables(connection)
377
378 if not connection.dialect.supports_alter:
379 from . import assertions
380
381 with assertions.expect_warnings(
382 "Can't sort tables", assert_=False
383 ):
384 metadata.drop_all(connection)
385 else:
386 metadata.drop_all(connection)
387
388 if not isinstance(engine_or_connection, Connection):
389 with engine_or_connection.begin() as connection:
390 go(connection)
391 else:
392 go(engine_or_connection)
393
394
395def drop_all_tables(
396 engine,
397 inspector,
398 schema=None,
399 consider_schemas=(None,),
400 include_names=None,
401):
402 if include_names is not None:
403 include_names = set(include_names)
404
405 if schema is not None:
406 assert consider_schemas == (
407 None,
408 ), "consider_schemas and schema are mutually exclusive"
409 consider_schemas = (schema,)
410
411 with engine.begin() as conn:
412 for table_key, fkcs in reversed(
413 inspector.sort_tables_on_foreign_key_dependency(
414 consider_schemas=consider_schemas
415 )
416 ):
417 if table_key:
418 if (
419 include_names is not None
420 and table_key[1] not in include_names
421 ):
422 continue
423 conn.execute(
424 DropTable(
425 Table(table_key[1], MetaData(), schema=table_key[0])
426 )
427 )
428 elif fkcs:
429 if not engine.dialect.supports_alter:
430 continue
431 for t_key, fkc in fkcs:
432 if (
433 include_names is not None
434 and t_key[1] not in include_names
435 ):
436 continue
437 tb = Table(
438 t_key[1],
439 MetaData(),
440 Column("x", Integer),
441 Column("y", Integer),
442 schema=t_key[0],
443 )
444 conn.execute(
445 DropConstraint(
446 ForeignKeyConstraint([tb.c.x], [tb.c.y], name=fkc)
447 )
448 )
449
450
451def teardown_events(event_cls):
452 @decorator
453 def decorate(fn, *arg, **kw):
454 try:
455 return fn(*arg, **kw)
456 finally:
457 event_cls._clear()
458
459 return decorate
460
461
462def total_size(o):
463 """Returns the approximate memory footprint an object and all of its
464 contents.
465
466 source: https://code.activestate.com/recipes/577504/
467
468
469 """
470
471 def dict_handler(d):
472 return chain.from_iterable(d.items())
473
474 all_handlers = {
475 tuple: iter,
476 list: iter,
477 deque: iter,
478 dict: dict_handler,
479 set: iter,
480 frozenset: iter,
481 }
482 seen = set() # track which object id's have already been seen
483 default_size = getsizeof(0) # estimate sizeof object without __sizeof__
484
485 def sizeof(o):
486 if id(o) in seen: # do not double count the same object
487 return 0
488 seen.add(id(o))
489 s = getsizeof(o, default_size)
490
491 for typ, handler in all_handlers.items():
492 if isinstance(o, typ):
493 s += sum(map(sizeof, handler(o)))
494 break
495 return s
496
497 return sizeof(o)
498
499
500def count_cache_key_tuples(tup):
501 """given a cache key tuple, counts how many instances of actual
502 tuples are found.
503
504 used to alert large jumps in cache key complexity.
505
506 """
507 stack = [tup]
508
509 sentinel = object()
510 num_elements = 0
511
512 while stack:
513 elem = stack.pop(0)
514 if elem is sentinel:
515 num_elements += 1
516 elif isinstance(elem, tuple):
517 if elem:
518 stack = list(elem) + [sentinel] + stack
519 return num_elements
520 