Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
assertions.py995 linesDownload Raw Back to testing
1# testing/assertions.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 defaultdict
13import contextlib
14from copy import copy
15from itertools import filterfalse
16import re
17import sys
18import warnings
19
20from . import assertsql
21from . import config
22from . import engines
23from . import mock
24from .exclusions import db_spec
25from .util import fail
26from .. import exc as sa_exc
27from .. import schema
28from .. import sql
29from .. import types as sqltypes
30from .. import util
31from ..engine import default
32from ..engine import url
33from ..sql.selectable import LABEL_STYLE_TABLENAME_PLUS_COL
34from ..util import decorator
35
36
37def expect_warnings(*messages, **kw):
38    """Context manager which expects one or more warnings.
39
40    With no arguments, squelches all SAWarning emitted via
41    sqlalchemy.util.warn and sqlalchemy.util.warn_limited.   Otherwise
42    pass string expressions that will match selected warnings via regex;
43    all non-matching warnings are sent through.
44
45    The expect version **asserts** that the warnings were in fact seen.
46
47    Note that the test suite sets SAWarning warnings to raise exceptions.
48
49    """  # noqa
50    return _expect_warnings_sqla_only(sa_exc.SAWarning, messages, **kw)
51
52
53@contextlib.contextmanager
54def expect_warnings_on(db, *messages, **kw):
55    """Context manager which expects one or more warnings on specific
56    dialects.
57
58    The expect version **asserts** that the warnings were in fact seen.
59
60    """
61    spec = db_spec(db)
62
63    if isinstance(db, str) and not spec(config._current):
64        yield
65    else:
66        with expect_warnings(*messages, **kw):
67            yield
68
69
70def emits_warning(*messages):
71    """Decorator form of expect_warnings().
72
73    Note that emits_warning does **not** assert that the warnings
74    were in fact seen.
75
76    """
77
78    @decorator
79    def decorate(fn, *args, **kw):
80        with expect_warnings(assert_=False, *messages):
81            return fn(*args, **kw)
82
83    return decorate
84
85
86def expect_deprecated(*messages, **kw):
87    return _expect_warnings_sqla_only(
88        sa_exc.SADeprecationWarning, messages, **kw
89    )
90
91
92def expect_deprecated_20(*messages, **kw):
93    return _expect_warnings_sqla_only(
94        sa_exc.Base20DeprecationWarning, messages, **kw
95    )
96
97
98def emits_warning_on(db, *messages):
99    """Mark a test as emitting a warning on a specific dialect.
100
101    With no arguments, squelches all SAWarning failures.  Or pass one or more
102    strings; these will be matched to the root of the warning description by
103    warnings.filterwarnings().
104
105    Note that emits_warning_on does **not** assert that the warnings
106    were in fact seen.
107
108    """
109
110    @decorator
111    def decorate(fn, *args, **kw):
112        with expect_warnings_on(db, assert_=False, *messages):
113            return fn(*args, **kw)
114
115    return decorate
116
117
118def uses_deprecated(*messages):
119    """Mark a test as immune from fatal deprecation warnings.
120
121    With no arguments, squelches all SADeprecationWarning failures.
122    Or pass one or more strings; these will be matched to the root
123    of the warning description by warnings.filterwarnings().
124
125    As a special case, you may pass a function name prefixed with //
126    and it will be re-written as needed to match the standard warning
127    verbiage emitted by the sqlalchemy.util.deprecated decorator.
128
129    Note that uses_deprecated does **not** assert that the warnings
130    were in fact seen.
131
132    """
133
134    @decorator
135    def decorate(fn, *args, **kw):
136        with expect_deprecated(*messages, assert_=False):
137            return fn(*args, **kw)
138
139    return decorate
140
141
142_FILTERS = None
143_SEEN = None
144_EXC_CLS = None
145
146
147def _expect_warnings_sqla_only(
148    exc_cls,
149    messages,
150    regex=True,
151    search_msg=False,
152    assert_=True,
153):
154    """SQLAlchemy internal use only _expect_warnings().
155
156    Alembic is using _expect_warnings() directly, and should be updated
157    to use this new interface.
158
159    """
160    return _expect_warnings(
161        exc_cls,
162        messages,
163        regex=regex,
164        search_msg=search_msg,
165        assert_=assert_,
166        raise_on_any_unexpected=True,
167    )
168
169
170@contextlib.contextmanager
171def _expect_warnings(
172    exc_cls,
173    messages,
174    regex=True,
175    search_msg=False,
176    assert_=True,
177    raise_on_any_unexpected=False,
178    squelch_other_warnings=False,
179):
180    global _FILTERS, _SEEN, _EXC_CLS
181
182    if regex or search_msg:
183        filters = [re.compile(msg, re.I | re.S) for msg in messages]
184    else:
185        filters = list(messages)
186
187    if _FILTERS is not None:
188        # nested call; update _FILTERS and _SEEN, return.  outer
189        # block will assert our messages
190        assert _SEEN is not None
191        assert _EXC_CLS is not None
192        _FILTERS.extend(filters)
193        _SEEN.update(filters)
194        _EXC_CLS += (exc_cls,)
195        yield
196    else:
197        seen = _SEEN = set(filters)
198        _FILTERS = filters
199        _EXC_CLS = (exc_cls,)
200
201        if raise_on_any_unexpected:
202
203            def real_warn(msg, *arg, **kw):
204                if isinstance(msg, sa_exc.SATestSuiteWarning):
205                    warnings.warn(msg, *arg, **kw)
206                else:
207                    raise AssertionError("Got unexpected warning: %r" % msg)
208
209        else:
210            real_warn = warnings.warn
211
212        def our_warn(msg, *arg, **kw):
213            if isinstance(msg, _EXC_CLS):
214                exception = type(msg)
215                msg = str(msg)
216            elif arg:
217                exception = arg[0]
218            else:
219                exception = None
220
221            if not exception or not issubclass(exception, _EXC_CLS):
222                if not squelch_other_warnings:
223                    return real_warn(msg, *arg, **kw)
224                else:
225                    return
226
227            if not filters and not raise_on_any_unexpected:
228                return
229
230            for filter_ in filters:
231                if (
232                    (search_msg and filter_.search(msg))
233                    or (regex and filter_.match(msg))
234                    or (not regex and filter_ == msg)
235                ):
236                    seen.discard(filter_)
237                    break
238            else:
239                if not squelch_other_warnings:
240                    real_warn(msg, *arg, **kw)
241
242        with mock.patch("warnings.warn", our_warn):
243            try:
244                yield
245            finally:
246                _SEEN = _FILTERS = _EXC_CLS = None
247
248                if assert_:
249                    assert not seen, "Warnings were not seen: %s" % ", ".join(
250                        "%r" % (s.pattern if regex else s) for s in seen
251                    )
252
253
254def global_cleanup_assertions():
255    """Check things that have to be finalized at the end of a test suite.
256
257    Hardcoded at the moment, a modular system can be built here
258    to support things like PG prepared transactions, tables all
259    dropped, etc.
260
261    """
262    _assert_no_stray_pool_connections()
263
264
265def _assert_no_stray_pool_connections():
266    engines.testing_reaper.assert_all_closed()
267
268
269def int_within_variance(expected, received, variance):
270    deviance = int(expected * variance)
271    assert (
272        abs(received - expected) < deviance
273    ), "Given int value %s is not within %d%% of expected value %s" % (
274        received,
275        variance * 100,
276        expected,
277    )
278
279
280def eq_regex(a, b, msg=None, flags=0):
281    assert re.match(b, a, flags), msg or "%r !~ %r" % (a, b)
282
283
284def eq_(a, b, msg=None):
285    """Assert a == b, with repr messaging on failure."""
286    assert a == b, msg or "%r != %r" % (a, b)
287
288
289def ne_(a, b, msg=None):
290    """Assert a != b, with repr messaging on failure."""
291    assert a != b, msg or "%r == %r" % (a, b)
292
293
294def le_(a, b, msg=None):
295    """Assert a <= b, with repr messaging on failure."""
296    assert a <= b, msg or "%r != %r" % (a, b)
297
298
299def is_instance_of(a, b, msg=None):
300    assert isinstance(a, b), msg or "%r is not an instance of %r" % (a, b)
301
302
303def is_none(a, msg=None):
304    is_(a, None, msg=msg)
305
306
307def is_not_none(a, msg=None):
308    is_not(a, None, msg=msg)
309
310
311def is_true(a, msg=None):
312    is_(bool(a), True, msg=msg)
313
314
315def is_false(a, msg=None):
316    is_(bool(a), False, msg=msg)
317
318
319def is_(a, b, msg=None):
320    """Assert a is b, with repr messaging on failure."""
321    assert a is b, msg or "%r is not %r" % (a, b)
322
323
324def is_not(a, b, msg=None):
325    """Assert a is not b, with repr messaging on failure."""
326    assert a is not b, msg or "%r is %r" % (a, b)
327
328
329# deprecated.  See #5429
330is_not_ = is_not
331
332
333def in_(a, b, msg=None):
334    """Assert a in b, with repr messaging on failure."""
335    assert a in b, msg or "%r not in %r" % (a, b)
336
337
338def not_in(a, b, msg=None):
339    """Assert a in not b, with repr messaging on failure."""
340    assert a not in b, msg or "%r is in %r" % (a, b)
341
342
343# deprecated.  See #5429
344not_in_ = not_in
345
346
347def startswith_(a, fragment, msg=None):
348    """Assert a.startswith(fragment), with repr messaging on failure."""
349    assert a.startswith(fragment), msg or "%r does not start with %r" % (
350        a,
351        fragment,
352    )
353
354
355def eq_ignore_whitespace(a, b, msg=None):
356    a = re.sub(r"^\s+?|\n", "", a)
357    a = re.sub(r" {2,}", " ", a)
358    a = re.sub(r"\t", "", a)
359    b = re.sub(r"^\s+?|\n", "", b)
360    b = re.sub(r" {2,}", " ", b)
361    b = re.sub(r"\t", "", b)
362
363    assert a == b, msg or "%r != %r" % (a, b)
364
365
366def _assert_proper_exception_context(exception):
367    """assert that any exception we're catching does not have a __context__
368    without a __cause__, and that __suppress_context__ is never set.
369
370    Python 3 will report nested as exceptions as "during the handling of
371    error X, error Y occurred". That's not what we want to do.  we want
372    these exceptions in a cause chain.
373
374    """
375
376    if (
377        exception.__context__ is not exception.__cause__
378        and not exception.__suppress_context__
379    ):
380        assert False, (
381            "Exception %r was correctly raised but did not set a cause, "
382            "within context %r as its cause."
383            % (exception, exception.__context__)
384        )
385
386
387def assert_raises(except_cls, callable_, *args, **kw):
388    return _assert_raises(except_cls, callable_, args, kw, check_context=True)
389
390
391def assert_raises_context_ok(except_cls, callable_, *args, **kw):
392    return _assert_raises(except_cls, callable_, args, kw)
393
394
395def assert_raises_message(except_cls, msg, callable_, *args, **kwargs):
396    return _assert_raises(
397        except_cls, callable_, args, kwargs, msg=msg, check_context=True
398    )
399
400
401def assert_warns(except_cls, callable_, *args, **kwargs):
402    """legacy adapter function for functions that were previously using
403    assert_raises with SAWarning or similar.
404
405    has some workarounds to accommodate the fact that the callable completes
406    with this approach rather than stopping at the exception raise.
407
408
409    """
410    with _expect_warnings_sqla_only(except_cls, [".*"]):
411        return callable_(*args, **kwargs)
412
413
414def assert_warns_message(except_cls, msg, callable_, *args, **kwargs):
415    """legacy adapter function for functions that were previously using
416    assert_raises with SAWarning or similar.
417
418    has some workarounds to accommodate the fact that the callable completes
419    with this approach rather than stopping at the exception raise.
420
421    Also uses regex.search() to match the given message to the error string
422    rather than regex.match().
423
424    """
425    with _expect_warnings_sqla_only(
426        except_cls,
427        [msg],
428        search_msg=True,
429        regex=False,
430    ):
431        return callable_(*args, **kwargs)
432
433
434def assert_raises_message_context_ok(
435    except_cls, msg, callable_, *args, **kwargs
436):
437    return _assert_raises(except_cls, callable_, args, kwargs, msg=msg)
438
439
440def _assert_raises(
441    except_cls, callable_, args, kwargs, msg=None, check_context=False
442):
443    with _expect_raises(except_cls, msg, check_context) as ec:
444        callable_(*args, **kwargs)
445    return ec.error
446
447
448class _ErrorContainer:
449    error = None
450
451
452@contextlib.contextmanager
453def _expect_raises(except_cls, msg=None, check_context=False):
454    if (
455        isinstance(except_cls, type)
456        and issubclass(except_cls, Warning)
457        or isinstance(except_cls, Warning)
458    ):
459        raise TypeError(
460            "Use expect_warnings for warnings, not "
461            "expect_raises / assert_raises"
462        )
463    ec = _ErrorContainer()
464    if check_context:
465        are_we_already_in_a_traceback = sys.exc_info()[0]
466    try:
467        yield ec
468        success = False
469    except except_cls as err:
470        ec.error = err
471        success = True
472        if msg is not None:
473            # I'm often pdbing here, and "err" above isn't
474            # in scope, so assign the string explicitly
475            error_as_string = str(err)
476            assert re.search(msg, error_as_string, re.UNICODE), "%r !~ %s" % (
477                msg,
478                error_as_string,
479            )
480        if check_context and not are_we_already_in_a_traceback:
481            _assert_proper_exception_context(err)
482        print(str(err).encode("utf-8"))
483
484    # it's generally a good idea to not carry traceback objects outside
485    # of the except: block, but in this case especially we seem to have
486    # hit some bug in either python 3.10.0b2 or greenlet or both which
487    # this seems to fix:
488    # https://github.com/python-greenlet/greenlet/issues/242
489    del ec
490
491    # assert outside the block so it works for AssertionError too !
492    assert success, "Callable did not raise an exception"
493
494
495def expect_raises(except_cls, check_context=True):
496    return _expect_raises(except_cls, check_context=check_context)
497
498
499def expect_raises_message(except_cls, msg, check_context=True):
500    return _expect_raises(except_cls, msg=msg, check_context=check_context)
501
502
503class AssertsCompiledSQL:
504    def assert_compile(
505        self,
506        clause,
507        result,
508        params=None,
509        checkparams=None,
510        for_executemany=False,
511        check_literal_execute=None,
512        check_post_param=None,
513        dialect=None,
514        checkpositional=None,
515        check_prefetch=None,
516        use_default_dialect=False,
517        allow_dialect_select=False,
518        supports_default_values=True,
519        supports_native_boolean=False,
520        supports_default_metavalue=True,
521        literal_binds=False,
522        render_postcompile=False,
523        schema_translate_map=None,
524        render_schema_translate=False,
525        default_schema_name=None,
526        from_linting=False,
527        check_param_order=True,
528        use_literal_execute_for_simple_int=False,
529    ):
530        if use_default_dialect:
531            dialect = default.DefaultDialect()
532            dialect.supports_default_values = supports_default_values
533            dialect.supports_default_metavalue = supports_default_metavalue
534            dialect.supports_native_boolean = supports_native_boolean
535        elif allow_dialect_select:
536            dialect = None
537        else:
538            if dialect is None:
539                dialect = getattr(self, "__dialect__", None)
540
541            if dialect is None:
542                dialect = config.db.dialect
543            elif dialect == "default" or dialect == "default_qmark":
544                if dialect == "default":
545                    dialect = default.DefaultDialect()
546                else:
547                    dialect = default.DefaultDialect("qmark")
548                dialect.supports_default_values = supports_default_values
549                dialect.supports_default_metavalue = supports_default_metavalue
550            elif dialect == "default_enhanced":
551                dialect = default.StrCompileDialect()
552            elif isinstance(dialect, str):
553                dialect = url.URL.create(dialect).get_dialect()()
554
555        if default_schema_name:
556            dialect.default_schema_name = default_schema_name
557
558        kw = {}
559        compile_kwargs = {}
560
561        if schema_translate_map:
562            kw["schema_translate_map"] = schema_translate_map
563
564        if params is not None:
565            kw["column_keys"] = list(params)
566
567        if literal_binds:
568            compile_kwargs["literal_binds"] = True
569
570        if render_postcompile:
571            compile_kwargs["render_postcompile"] = True
572
573        if use_literal_execute_for_simple_int:
574            compile_kwargs["use_literal_execute_for_simple_int"] = True
575
576        if for_executemany:
577            kw["for_executemany"] = True
578
579        if render_schema_translate:
580            kw["render_schema_translate"] = True
581
582        if from_linting or getattr(self, "assert_from_linting", False):
583            kw["linting"] = sql.FROM_LINTING
584
585        from sqlalchemy import orm
586
587        if isinstance(clause, orm.Query):
588            stmt = clause._statement_20()
589            stmt._label_style = LABEL_STYLE_TABLENAME_PLUS_COL
590            clause = stmt
591
592        if compile_kwargs:
593            kw["compile_kwargs"] = compile_kwargs
594
595        class DontAccess:
596            def __getattribute__(self, key):
597                raise NotImplementedError(
598                    "compiler accessed .statement; use "
599                    "compiler.current_executable"
600                )
601
602        class CheckCompilerAccess:
603            def __init__(self, test_statement):
604                self.test_statement = test_statement
605                self._annotations = {}
606                self.supports_execution = getattr(
607                    test_statement, "supports_execution", False
608                )
609
610                if self.supports_execution:
611                    self._execution_options = test_statement._execution_options
612
613                    if hasattr(test_statement, "_returning"):
614                        self._returning = test_statement._returning
615                    if hasattr(test_statement, "_inline"):
616                        self._inline = test_statement._inline
617                    if hasattr(test_statement, "_return_defaults"):
618                        self._return_defaults = test_statement._return_defaults
619
620            @property
621            def _variant_mapping(self):
622                return self.test_statement._variant_mapping
623
624            def _default_dialect(self):
625                return self.test_statement._default_dialect()
626
627            def compile(self, dialect, **kw):
628                return self.test_statement.compile.__func__(
629                    self, dialect=dialect, **kw
630                )
631
632            def _compiler(self, dialect, **kw):
633                return self.test_statement._compiler.__func__(
634                    self, dialect, **kw
635                )
636
637            def _compiler_dispatch(self, compiler, **kwargs):
638                if hasattr(compiler, "statement"):
639                    with mock.patch.object(
640                        compiler, "statement", DontAccess()
641                    ):
642                        return self.test_statement._compiler_dispatch(
643                            compiler, **kwargs
644                        )
645                else:
646                    return self.test_statement._compiler_dispatch(
647                        compiler, **kwargs
648                    )
649
650        # no construct can assume it's the "top level" construct in all cases
651        # as anything can be nested.  ensure constructs don't assume they
652        # are the "self.statement" element
653        c = CheckCompilerAccess(clause).compile(dialect=dialect, **kw)
654
655        if isinstance(clause, sqltypes.TypeEngine):
656            cache_key_no_warnings = clause._static_cache_key
657            if cache_key_no_warnings:
658                hash(cache_key_no_warnings)
659        else:
660            cache_key_no_warnings = clause._generate_cache_key()
661            if cache_key_no_warnings:
662                hash(cache_key_no_warnings[0])
663
664        param_str = repr(getattr(c, "params", {}))
665        param_str = param_str.encode("utf-8").decode("ascii", "ignore")
666        print(("\nSQL String:\n" + str(c) + param_str).encode("utf-8"))
667
668        cc = re.sub(r"[\n\t]", "", str(c))
669
670        eq_(cc, result, "%r != %r on dialect %r" % (cc, result, dialect))
671
672        if checkparams is not None:
673            if render_postcompile:
674                expanded_state = c.construct_expanded_state(
675                    params, escape_names=False
676                )
677                eq_(expanded_state.parameters, checkparams)
678            else:
679                eq_(c.construct_params(params), checkparams)
680        if checkpositional is not None:
681            if render_postcompile:
682                expanded_state = c.construct_expanded_state(
683                    params, escape_names=False
684                )
685                eq_(
686                    tuple(
687                        [
688                            expanded_state.parameters[x]
689                            for x in expanded_state.positiontup
690                        ]
691                    ),
692                    checkpositional,
693                )
694            else:
695                p = c.construct_params(params, escape_names=False)
696                eq_(tuple([p[x] for x in c.positiontup]), checkpositional)
697        if check_prefetch is not None:
698            eq_(c.prefetch, check_prefetch)
699        if check_literal_execute is not None:
700            eq_(
701                {
702                    c.bind_names[b]: b.effective_value
703                    for b in c.literal_execute_params
704                },
705                check_literal_execute,
706            )
707        if check_post_param is not None:
708            eq_(
709                {
710                    c.bind_names[b]: b.effective_value
711                    for b in c.post_compile_params
712                },
713                check_post_param,
714            )
715        if check_param_order and getattr(c, "params", None):
716
717            def get_dialect(paramstyle, positional):
718                cp = copy(dialect)
719                cp.paramstyle = paramstyle
720                cp.positional = positional
721                return cp
722
723            pyformat_dialect = get_dialect("pyformat", False)
724            pyformat_c = clause.compile(dialect=pyformat_dialect, **kw)
725            stmt = re.sub(r"[\n\t]", "", str(pyformat_c))
726
727            qmark_dialect = get_dialect("qmark", True)
728            qmark_c = clause.compile(dialect=qmark_dialect, **kw)
729            values = list(qmark_c.positiontup)
730            escaped = qmark_c.escaped_bind_names
731
732            for post_param in (
733                qmark_c.post_compile_params | qmark_c.literal_execute_params
734            ):
735                name = qmark_c.bind_names[post_param]
736                if name in values:
737                    values = [v for v in values if v != name]
738            positions = []
739            pos_by_value = defaultdict(list)
740            for v in values:
741                try:
742                    if v in pos_by_value:
743                        start = pos_by_value[v][-1]
744                    else:
745                        start = 0
746                    esc = escaped.get(v, v)
747                    pos = stmt.index("%%(%s)s" % (esc,), start) + 2
748                    positions.append(pos)
749                    pos_by_value[v].append(pos)
750                except ValueError:
751                    msg = "Expected to find bindparam %r in %r" % (v, stmt)
752                    assert False, msg
753
754            ordered = all(
755                positions[i - 1] < positions[i]
756                for i in range(1, len(positions))
757            )
758
759            expected = [v for _, v in sorted(zip(positions, values))]
760
761            msg = (
762                "Order of parameters %s does not match the order "
763                "in the statement %s. Statement %r" % (values, expected, stmt)
764            )
765
766            is_true(ordered, msg)
767
768
769class ComparesTables:
770    def assert_tables_equal(
771        self,
772        table,
773        reflected_table,
774        strict_types=False,
775        strict_constraints=True,
776    ):
777        assert len(table.c) == len(reflected_table.c)
778        for c, reflected_c in zip(table.c, reflected_table.c):
779            eq_(c.name, reflected_c.name)
780            assert reflected_c is reflected_table.c[c.name]
781
782            if strict_constraints:
783                eq_(c.primary_key, reflected_c.primary_key)
784                eq_(c.nullable, reflected_c.nullable)
785
786            if strict_types:
787                msg = "Type '%s' doesn't correspond to type '%s'"
788                assert isinstance(reflected_c.type, type(c.type)), msg % (
789                    reflected_c.type,
790                    c.type,
791                )
792            else:
793                self.assert_types_base(reflected_c, c)
794
795            if isinstance(c.type, sqltypes.String):
796                eq_(c.type.length, reflected_c.type.length)
797
798            if strict_constraints:
799                eq_(
800                    {f.column.name for f in c.foreign_keys},
801                    {f.column.name for f in reflected_c.foreign_keys},
802                )
803            if c.server_default:
804                assert isinstance(
805                    reflected_c.server_default, schema.FetchedValue
806                )
807
808        if strict_constraints:
809            assert len(table.primary_key) == len(reflected_table.primary_key)
810            for c in table.primary_key:
811                assert reflected_table.primary_key.columns[c.name] is not None
812
813    def assert_types_base(self, c1, c2):
814        assert c1.type._compare_type_affinity(
815            c2.type
816        ), "On column %r, type '%s' doesn't correspond to type '%s'" % (
817            c1.name,
818            c1.type,
819            c2.type,
820        )
821
822
823class AssertsExecutionResults:
824    def assert_result(self, result, class_, *objects):
825        result = list(result)
826        print(repr(result))
827        self.assert_list(result, class_, objects)
828
829    def assert_list(self, result, class_, list_):
830        self.assert_(
831            len(result) == len(list_),
832            "result list is not the same size as test list, "
833            + "for class "
834            + class_.__name__,
835        )
836        for i in range(0, len(list_)):
837            self.assert_row(class_, result[i], list_[i])
838
839    def assert_row(self, class_, rowobj, desc):
840        self.assert_(
841            rowobj.__class__ is class_, "item class is not " + repr(class_)
842        )
843        for key, value in desc.items():
844            if isinstance(value, tuple):
845                if isinstance(value[1], list):
846                    self.assert_list(getattr(rowobj, key), value[0], value[1])
847                else:
848                    self.assert_row(value[0], getattr(rowobj, key), value[1])
849            else:
850                self.assert_(
851                    getattr(rowobj, key) == value,
852                    "attribute %s value %s does not match %s"
853                    % (key, getattr(rowobj, key), value),
854                )
855
856    def assert_unordered_result(self, result, cls, *expected):
857        """As assert_result, but the order of objects is not considered.
858
859        The algorithm is very expensive but not a big deal for the small
860        numbers of rows that the test suite manipulates.
861        """
862
863        class immutabledict(dict):
864            def __hash__(self):
865                return id(self)
866
867        found = util.IdentitySet(result)
868        expected = {immutabledict(e) for e in expected}
869
870        for wrong in filterfalse(lambda o: isinstance(o, cls), found):
871            fail(
872                'Unexpected type "%s", expected "%s"'
873                % (type(wrong).__name__, cls.__name__)
874            )
875
876        if len(found) != len(expected):
877            fail(
878                'Unexpected object count "%s", expected "%s"'
879                % (len(found), len(expected))
880            )
881
882        NOVALUE = object()
883
884        def _compare_item(obj, spec):
885            for key, value in spec.items():
886                if isinstance(value, tuple):
887                    try:
888                        self.assert_unordered_result(
889                            getattr(obj, key), value[0], *value[1]
890                        )
891                    except AssertionError:
892                        return False
893                else:
894                    if getattr(obj, key, NOVALUE) != value:
895                        return False
896            return True
897
898        for expected_item in expected:
899            for found_item in found:
900                if _compare_item(found_item, expected_item):
901                    found.remove(found_item)
902                    break
903            else:
904                fail(
905                    "Expected %s instance with attributes %s not found."
906                    % (cls.__name__, repr(expected_item))
907                )
908        return True
909
910    def sql_execution_asserter(self, db=None):
911        if db is None:
912            from . import db as db
913
914        return assertsql.assert_engine(db)
915
916    def assert_sql_execution(self, db, callable_, *rules):
917        with self.sql_execution_asserter(db) as asserter:
918            result = callable_()
919        asserter.assert_(*rules)
920        return result
921
922    def assert_sql(self, db, callable_, rules):
923        newrules = []
924        for rule in rules:
925            if isinstance(rule, dict):
926                newrule = assertsql.AllOf(
927                    *[assertsql.CompiledSQL(k, v) for k, v in rule.items()]
928                )
929            else:
930                newrule = assertsql.CompiledSQL(*rule)
931            newrules.append(newrule)
932
933        return self.assert_sql_execution(db, callable_, *newrules)
934
935    def assert_sql_count(self, db, callable_, count):
936        return self.assert_sql_execution(
937            db, callable_, assertsql.CountStatements(count)
938        )
939
940    @contextlib.contextmanager
941    def assert_execution(self, db, *rules):
942        with self.sql_execution_asserter(db) as asserter:
943            yield
944        asserter.assert_(*rules)
945
946    def assert_statement_count(self, db, count):
947        return self.assert_execution(db, assertsql.CountStatements(count))
948
949    @contextlib.contextmanager
950    def assert_statement_count_multi_db(self, dbs, counts):
951        recs = [
952            (self.sql_execution_asserter(db), db, count)
953            for (db, count) in zip(dbs, counts)
954        ]
955        asserters = []
956        for ctx, db, count in recs:
957            asserters.append(ctx.__enter__())
958        try:
959            yield
960        finally:
961            for asserter, (ctx, db, count) in zip(asserters, recs):
962                ctx.__exit__(None, None, None)
963                asserter.assert_(assertsql.CountStatements(count))
964
965
966class ComparesIndexes:
967    def compare_table_index_with_expected(
968        self, table: schema.Table, expected: list, dialect_name: str
969    ):
970        eq_(len(table.indexes), len(expected))
971        idx_dict = {idx.name: idx for idx in table.indexes}
972        for exp in expected:
973            idx = idx_dict[exp["name"]]
974            eq_(idx.unique, exp["unique"])
975            cols = [c for c in exp["column_names"] if c is not None]
976            eq_(len(idx.columns), len(cols))
977            for c in cols:
978                is_true(c in idx.columns)
979            exprs = exp.get("expressions")
980            if exprs:
981                eq_(len(idx.expressions), len(exprs))
982                for idx_exp, expr, col in zip(
983                    idx.expressions, exprs, exp["column_names"]
984                ):
985                    if col is None:
986                        eq_(idx_exp.text, expr)
987            if (
988                exp.get("dialect_options")
989                and f"{dialect_name}_include" in exp["dialect_options"]
990            ):
991                eq_(
992                    idx.dialect_options[dialect_name]["include"],
993                    exp["dialect_options"][f"{dialect_name}_include"],
994                )
995 
codekingpro/portable-devtools · Team Ai