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