Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
assertsql.py521 linesDownload Raw Back to testing
1# testing/assertsql.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
12import collections
13import contextlib
14import itertools
15import re
16
17from .. import event
18from ..engine import url
19from ..engine.default import DefaultDialect
20from ..schema import BaseDDLElement
21
22
23class AssertRule:
24    is_consumed = False
25    errormessage = None
26    consume_statement = True
27
28    def process_statement(self, execute_observed):
29        pass
30
31    def no_more_statements(self):
32        assert False, (
33            "All statements are complete, but pending "
34            "assertion rules remain"
35        )
36
37
38class SQLMatchRule(AssertRule):
39    pass
40
41
42class CursorSQL(SQLMatchRule):
43    def __init__(self, statement, params=None, consume_statement=True):
44        self.statement = statement
45        self.params = params
46        self.consume_statement = consume_statement
47
48    def process_statement(self, execute_observed):
49        stmt = execute_observed.statements[0]
50        if self.statement != stmt.statement or (
51            self.params is not None and self.params != stmt.parameters
52        ):
53            self.consume_statement = True
54            self.errormessage = (
55                "Testing for exact SQL %s parameters %s received %s %s"
56                % (
57                    self.statement,
58                    self.params,
59                    stmt.statement,
60                    stmt.parameters,
61                )
62            )
63        else:
64            execute_observed.statements.pop(0)
65            self.is_consumed = True
66            if not execute_observed.statements:
67                self.consume_statement = True
68
69
70class CompiledSQL(SQLMatchRule):
71    def __init__(
72        self, statement, params=None, dialect="default", enable_returning=True
73    ):
74        self.statement = statement
75        self.params = params
76        self.dialect = dialect
77        self.enable_returning = enable_returning
78
79    def _compare_sql(self, execute_observed, received_statement):
80        stmt = re.sub(r"[\n\t]", "", self.statement)
81        return received_statement == stmt
82
83    def _compile_dialect(self, execute_observed):
84        if self.dialect == "default":
85            dialect = DefaultDialect()
86            # this is currently what tests are expecting
87            # dialect.supports_default_values = True
88            dialect.supports_default_metavalue = True
89
90            if self.enable_returning:
91                dialect.insert_returning = dialect.update_returning = (
92                    dialect.delete_returning
93                ) = True
94                dialect.use_insertmanyvalues = True
95                dialect.supports_multivalues_insert = True
96                dialect.update_returning_multifrom = True
97                dialect.delete_returning_multifrom = True
98                # dialect.favor_returning_over_lastrowid = True
99                # dialect.insert_null_pk_still_autoincrements = True
100
101                # this is calculated but we need it to be True for this
102                # to look like all the current RETURNING dialects
103                assert dialect.insert_executemany_returning
104
105            return dialect
106        else:
107            return url.URL.create(self.dialect).get_dialect()()
108
109    def _received_statement(self, execute_observed):
110        """reconstruct the statement and params in terms
111        of a target dialect, which for CompiledSQL is just DefaultDialect."""
112
113        context = execute_observed.context
114        compare_dialect = self._compile_dialect(execute_observed)
115
116        # received_statement runs a full compile().  we should not need to
117        # consider extracted_parameters; if we do this indicates some state
118        # is being sent from a previous cached query, which some misbehaviors
119        # in the ORM can cause, see #6881
120        cache_key = None  # execute_observed.context.compiled.cache_key
121        extracted_parameters = (
122            None  # execute_observed.context.extracted_parameters
123        )
124
125        if "schema_translate_map" in context.execution_options:
126            map_ = context.execution_options["schema_translate_map"]
127        else:
128            map_ = None
129
130        if isinstance(execute_observed.clauseelement, BaseDDLElement):
131            compiled = execute_observed.clauseelement.compile(
132                dialect=compare_dialect,
133                schema_translate_map=map_,
134            )
135        else:
136            compiled = execute_observed.clauseelement.compile(
137                cache_key=cache_key,
138                dialect=compare_dialect,
139                column_keys=context.compiled.column_keys,
140                for_executemany=context.compiled.for_executemany,
141                schema_translate_map=map_,
142            )
143        _received_statement = re.sub(r"[\n\t]", "", str(compiled))
144        parameters = execute_observed.parameters
145
146        if not parameters:
147            _received_parameters = [
148                compiled.construct_params(
149                    extracted_parameters=extracted_parameters
150                )
151            ]
152        else:
153            _received_parameters = [
154                compiled.construct_params(
155                    m, extracted_parameters=extracted_parameters
156                )
157                for m in parameters
158            ]
159
160        return _received_statement, _received_parameters
161
162    def process_statement(self, execute_observed):
163        context = execute_observed.context
164
165        _received_statement, _received_parameters = self._received_statement(
166            execute_observed
167        )
168        params = self._all_params(context)
169
170        equivalent = self._compare_sql(execute_observed, _received_statement)
171
172        if equivalent:
173            if params is not None:
174                all_params = list(params)
175                all_received = list(_received_parameters)
176                while all_params and all_received:
177                    param = dict(all_params.pop(0))
178
179                    for idx, received in enumerate(list(all_received)):
180                        # do a positive compare only
181                        for param_key in param:
182                            # a key in param did not match current
183                            # 'received'
184                            if (
185                                param_key not in received
186                                or received[param_key] != param[param_key]
187                            ):
188                                break
189                        else:
190                            # all keys in param matched 'received';
191                            # onto next param
192                            del all_received[idx]
193                            break
194                    else:
195                        # param did not match any entry
196                        # in all_received
197                        equivalent = False
198                        break
199                if all_params or all_received:
200                    equivalent = False
201
202        if equivalent:
203            self.is_consumed = True
204            self.errormessage = None
205        else:
206            self.errormessage = self._failure_message(
207                execute_observed, params
208            ) % {
209                "received_statement": _received_statement,
210                "received_parameters": _received_parameters,
211            }
212
213    def _all_params(self, context):
214        if self.params:
215            if callable(self.params):
216                params = self.params(context)
217            else:
218                params = self.params
219            if not isinstance(params, list):
220                params = [params]
221            return params
222        else:
223            return None
224
225    def _failure_message(self, execute_observed, expected_params):
226        return (
227            "Testing for compiled statement\n%r partial params %s, "
228            "received\n%%(received_statement)r with params "
229            "%%(received_parameters)r"
230            % (
231                self.statement.replace("%", "%%"),
232                repr(expected_params).replace("%", "%%"),
233            )
234        )
235
236
237class RegexSQL(CompiledSQL):
238    def __init__(
239        self, regex, params=None, dialect="default", enable_returning=False
240    ):
241        SQLMatchRule.__init__(self)
242        self.regex = re.compile(regex)
243        self.orig_regex = regex
244        self.params = params
245        self.dialect = dialect
246        self.enable_returning = enable_returning
247
248    def _failure_message(self, execute_observed, expected_params):
249        return (
250            "Testing for compiled statement ~%r partial params %s, "
251            "received %%(received_statement)r with params "
252            "%%(received_parameters)r"
253            % (
254                self.orig_regex.replace("%", "%%"),
255                repr(expected_params).replace("%", "%%"),
256            )
257        )
258
259    def _compare_sql(self, execute_observed, received_statement):
260        return bool(self.regex.match(received_statement))
261
262
263class DialectSQL(CompiledSQL):
264    def _compile_dialect(self, execute_observed):
265        return execute_observed.context.dialect
266
267    def _compare_no_space(self, real_stmt, received_stmt):
268        stmt = re.sub(r"[\n\t]", "", real_stmt)
269        return received_stmt == stmt
270
271    def _received_statement(self, execute_observed):
272        received_stmt, received_params = super()._received_statement(
273            execute_observed
274        )
275
276        # TODO: why do we need this part?
277        for real_stmt in execute_observed.statements:
278            if self._compare_no_space(
279                real_stmt.context.statement, received_stmt
280            ):
281                break
282        else:
283            raise AssertionError(
284                "Can't locate compiled statement %r in list of "
285                "statements actually invoked" % received_stmt
286            )
287
288        return received_stmt, execute_observed.context.compiled_parameters
289
290    def _dialect_adjusted_statement(self, dialect):
291        paramstyle = dialect.paramstyle
292        stmt = re.sub(r"[\n\t]", "", self.statement)
293
294        # temporarily escape out PG double colons
295        stmt = stmt.replace("::", "!!")
296
297        if paramstyle == "pyformat":
298            stmt = re.sub(r":([\w_]+)", r"%(\1)s", stmt)
299        else:
300            # positional params
301            repl = None
302            if paramstyle == "qmark":
303                repl = "?"
304            elif paramstyle == "format":
305                repl = r"%s"
306            elif paramstyle.startswith("numeric"):
307                counter = itertools.count(1)
308
309                num_identifier = "$" if paramstyle == "numeric_dollar" else ":"
310
311                def repl(m):
312                    return f"{num_identifier}{next(counter)}"
313
314            stmt = re.sub(r":([\w_]+)", repl, stmt)
315
316        # put them back
317        stmt = stmt.replace("!!", "::")
318
319        return stmt
320
321    def _compare_sql(self, execute_observed, received_statement):
322        stmt = self._dialect_adjusted_statement(
323            execute_observed.context.dialect
324        )
325        return received_statement == stmt
326
327    def _failure_message(self, execute_observed, expected_params):
328        return (
329            "Testing for compiled statement\n%r partial params %s, "
330            "received\n%%(received_statement)r with params "
331            "%%(received_parameters)r"
332            % (
333                self._dialect_adjusted_statement(
334                    execute_observed.context.dialect
335                ).replace("%", "%%"),
336                repr(expected_params).replace("%", "%%"),
337            )
338        )
339
340
341class CountStatements(AssertRule):
342    def __init__(self, count):
343        self.count = count
344        self._statement_count = 0
345
346    def process_statement(self, execute_observed):
347        self._statement_count += 1
348
349    def no_more_statements(self):
350        if self.count != self._statement_count:
351            assert False, "desired statement count %d does not match %d" % (
352                self.count,
353                self._statement_count,
354            )
355
356
357class AllOf(AssertRule):
358    def __init__(self, *rules):
359        self.rules = set(rules)
360
361    def process_statement(self, execute_observed):
362        for rule in list(self.rules):
363            rule.errormessage = None
364            rule.process_statement(execute_observed)
365            if rule.is_consumed:
366                self.rules.discard(rule)
367                if not self.rules:
368                    self.is_consumed = True
369                break
370            elif not rule.errormessage:
371                # rule is not done yet
372                self.errormessage = None
373                break
374        else:
375            self.errormessage = list(self.rules)[0].errormessage
376
377
378class EachOf(AssertRule):
379    def __init__(self, *rules):
380        self.rules = list(rules)
381
382    def process_statement(self, execute_observed):
383        if not self.rules:
384            self.is_consumed = True
385            self.consume_statement = False
386
387        while self.rules:
388            rule = self.rules[0]
389            rule.process_statement(execute_observed)
390            if rule.is_consumed:
391                self.rules.pop(0)
392            elif rule.errormessage:
393                self.errormessage = rule.errormessage
394            if rule.consume_statement:
395                break
396
397        if not self.rules:
398            self.is_consumed = True
399
400    def no_more_statements(self):
401        if self.rules and not self.rules[0].is_consumed:
402            self.rules[0].no_more_statements()
403        elif self.rules:
404            super().no_more_statements()
405
406
407class Conditional(EachOf):
408    def __init__(self, condition, rules, else_rules):
409        if condition:
410            super().__init__(*rules)
411        else:
412            super().__init__(*else_rules)
413
414
415class Or(AllOf):
416    def process_statement(self, execute_observed):
417        for rule in self.rules:
418            rule.process_statement(execute_observed)
419            if rule.is_consumed:
420                self.is_consumed = True
421                break
422        else:
423            self.errormessage = list(self.rules)[0].errormessage
424
425
426class SQLExecuteObserved:
427    def __init__(self, context, clauseelement, multiparams, params):
428        self.context = context
429        self.clauseelement = clauseelement
430
431        if multiparams:
432            self.parameters = multiparams
433        elif params:
434            self.parameters = [params]
435        else:
436            self.parameters = []
437        self.statements = []
438
439    def __repr__(self):
440        return str(self.statements)
441
442
443class SQLCursorExecuteObserved(
444    collections.namedtuple(
445        "SQLCursorExecuteObserved",
446        ["statement", "parameters", "context", "executemany"],
447    )
448):
449    pass
450
451
452class SQLAsserter:
453    def __init__(self):
454        self.accumulated = []
455
456    def _close(self):
457        self._final = self.accumulated
458        del self.accumulated
459
460    def assert_(self, *rules):
461        rule = EachOf(*rules)
462
463        observed = list(self._final)
464        while observed:
465            statement = observed.pop(0)
466            rule.process_statement(statement)
467            if rule.is_consumed:
468                break
469            elif rule.errormessage:
470                assert False, rule.errormessage
471        if observed:
472            assert False, "Additional SQL statements remain:\n%s" % observed
473        elif not rule.is_consumed:
474            rule.no_more_statements()
475
476
477@contextlib.contextmanager
478def assert_engine(engine):
479    asserter = SQLAsserter()
480
481    orig = []
482
483    @event.listens_for(engine, "before_execute")
484    def connection_execute(
485        conn, clauseelement, multiparams, params, execution_options
486    ):
487        conn._WORKAROUND_ISSUE_13018 = True
488        # grab the original statement + params before any cursor
489        # execution
490        orig[:] = clauseelement, multiparams, params
491
492    @event.listens_for(engine, "after_cursor_execute")
493    def cursor_execute(
494        conn, cursor, statement, parameters, context, executemany
495    ):
496        if not context:
497            return
498        # then grab real cursor statements and associate them all
499        # around a single context
500        if (
501            asserter.accumulated
502            and asserter.accumulated[-1].context is context
503        ):
504            obs = asserter.accumulated[-1]
505        else:
506            obs = SQLExecuteObserved(context, orig[0], orig[1], orig[2])
507            asserter.accumulated.append(obs)
508
509        obs.statements.append(
510            SQLCursorExecuteObserved(
511                statement, parameters, context, executemany
512            )
513        )
514
515    try:
516        yield asserter
517    finally:
518        event.remove(engine, "after_cursor_execute", cursor_execute)
519        event.remove(engine, "before_execute", connection_execute)
520        asserter._close()
521