codekingpro/portable-devtools
115k
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 