Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
assertsql.py517 linesDownload Raw Back to testing
1# testing/assertsql.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
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(real_stmt.statement, received_stmt):
279                break
280        else:
281            raise AssertionError(
282                "Can't locate compiled statement %r in list of "
283                "statements actually invoked" % received_stmt
284            )
285
286        return received_stmt, execute_observed.context.compiled_parameters
287
288    def _dialect_adjusted_statement(self, dialect):
289        paramstyle = dialect.paramstyle
290        stmt = re.sub(r"[\n\t]", "", self.statement)
291
292        # temporarily escape out PG double colons
293        stmt = stmt.replace("::", "!!")
294
295        if paramstyle == "pyformat":
296            stmt = re.sub(r":([\w_]+)", r"%(\1)s", stmt)
297        else:
298            # positional params
299            repl = None
300            if paramstyle == "qmark":
301                repl = "?"
302            elif paramstyle == "format":
303                repl = r"%s"
304            elif paramstyle.startswith("numeric"):
305                counter = itertools.count(1)
306
307                num_identifier = "$" if paramstyle == "numeric_dollar" else ":"
308
309                def repl(m):
310                    return f"{num_identifier}{next(counter)}"
311
312            stmt = re.sub(r":([\w_]+)", repl, stmt)
313
314        # put them back
315        stmt = stmt.replace("!!", "::")
316
317        return stmt
318
319    def _compare_sql(self, execute_observed, received_statement):
320        stmt = self._dialect_adjusted_statement(
321            execute_observed.context.dialect
322        )
323        return received_statement == stmt
324
325    def _failure_message(self, execute_observed, expected_params):
326        return (
327            "Testing for compiled statement\n%r partial params %s, "
328            "received\n%%(received_statement)r with params "
329            "%%(received_parameters)r"
330            % (
331                self._dialect_adjusted_statement(
332                    execute_observed.context.dialect
333                ).replace("%", "%%"),
334                repr(expected_params).replace("%", "%%"),
335            )
336        )
337
338
339class CountStatements(AssertRule):
340    def __init__(self, count):
341        self.count = count
342        self._statement_count = 0
343
344    def process_statement(self, execute_observed):
345        self._statement_count += 1
346
347    def no_more_statements(self):
348        if self.count != self._statement_count:
349            assert False, "desired statement count %d does not match %d" % (
350                self.count,
351                self._statement_count,
352            )
353
354
355class AllOf(AssertRule):
356    def __init__(self, *rules):
357        self.rules = set(rules)
358
359    def process_statement(self, execute_observed):
360        for rule in list(self.rules):
361            rule.errormessage = None
362            rule.process_statement(execute_observed)
363            if rule.is_consumed:
364                self.rules.discard(rule)
365                if not self.rules:
366                    self.is_consumed = True
367                break
368            elif not rule.errormessage:
369                # rule is not done yet
370                self.errormessage = None
371                break
372        else:
373            self.errormessage = list(self.rules)[0].errormessage
374
375
376class EachOf(AssertRule):
377    def __init__(self, *rules):
378        self.rules = list(rules)
379
380    def process_statement(self, execute_observed):
381        if not self.rules:
382            self.is_consumed = True
383            self.consume_statement = False
384
385        while self.rules:
386            rule = self.rules[0]
387            rule.process_statement(execute_observed)
388            if rule.is_consumed:
389                self.rules.pop(0)
390            elif rule.errormessage:
391                self.errormessage = rule.errormessage
392            if rule.consume_statement:
393                break
394
395        if not self.rules:
396            self.is_consumed = True
397
398    def no_more_statements(self):
399        if self.rules and not self.rules[0].is_consumed:
400            self.rules[0].no_more_statements()
401        elif self.rules:
402            super().no_more_statements()
403
404
405class Conditional(EachOf):
406    def __init__(self, condition, rules, else_rules):
407        if condition:
408            super().__init__(*rules)
409        else:
410            super().__init__(*else_rules)
411
412
413class Or(AllOf):
414    def process_statement(self, execute_observed):
415        for rule in self.rules:
416            rule.process_statement(execute_observed)
417            if rule.is_consumed:
418                self.is_consumed = True
419                break
420        else:
421            self.errormessage = list(self.rules)[0].errormessage
422
423
424class SQLExecuteObserved:
425    def __init__(self, context, clauseelement, multiparams, params):
426        self.context = context
427        self.clauseelement = clauseelement
428
429        if multiparams:
430            self.parameters = multiparams
431        elif params:
432            self.parameters = [params]
433        else:
434            self.parameters = []
435        self.statements = []
436
437    def __repr__(self):
438        return str(self.statements)
439
440
441class SQLCursorExecuteObserved(
442    collections.namedtuple(
443        "SQLCursorExecuteObserved",
444        ["statement", "parameters", "context", "executemany"],
445    )
446):
447    pass
448
449
450class SQLAsserter:
451    def __init__(self):
452        self.accumulated = []
453
454    def _close(self):
455        self._final = self.accumulated
456        del self.accumulated
457
458    def assert_(self, *rules):
459        rule = EachOf(*rules)
460
461        observed = list(self._final)
462        while observed:
463            statement = observed.pop(0)
464            rule.process_statement(statement)
465            if rule.is_consumed:
466                break
467            elif rule.errormessage:
468                assert False, rule.errormessage
469        if observed:
470            assert False, "Additional SQL statements remain:\n%s" % observed
471        elif not rule.is_consumed:
472            rule.no_more_statements()
473
474
475@contextlib.contextmanager
476def assert_engine(engine):
477    asserter = SQLAsserter()
478
479    orig = []
480
481    @event.listens_for(engine, "before_execute")
482    def connection_execute(
483        conn, clauseelement, multiparams, params, execution_options
484    ):
485        # grab the original statement + params before any cursor
486        # execution
487        orig[:] = clauseelement, multiparams, params
488
489    @event.listens_for(engine, "after_cursor_execute")
490    def cursor_execute(
491        conn, cursor, statement, parameters, context, executemany
492    ):
493        if not context:
494            return
495        # then grab real cursor statements and associate them all
496        # around a single context
497        if (
498            asserter.accumulated
499            and asserter.accumulated[-1].context is context
500        ):
501            obs = asserter.accumulated[-1]
502        else:
503            obs = SQLExecuteObserved(context, orig[0], orig[1], orig[2])
504            asserter.accumulated.append(obs)
505        obs.statements.append(
506            SQLCursorExecuteObserved(
507                statement, parameters, context, executemany
508            )
509        )
510
511    try:
512        yield asserter
513    finally:
514        event.remove(engine, "after_cursor_execute", cursor_execute)
515        event.remove(engine, "before_execute", connection_execute)
516        asserter._close()
517 
codekingpro/portable-devtools · Team Ai