Team Ai
Datasetpublic

codekingpro/portable-devtools

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