SQLAlchemy 2.0.36__cp313-cp313-win32.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (273) hide show
  1. SQLAlchemy-2.0.36.dist-info/LICENSE +19 -0
  2. SQLAlchemy-2.0.36.dist-info/METADATA +243 -0
  3. SQLAlchemy-2.0.36.dist-info/RECORD +273 -0
  4. SQLAlchemy-2.0.36.dist-info/WHEEL +5 -0
  5. SQLAlchemy-2.0.36.dist-info/top_level.txt +1 -0
  6. sqlalchemy/__init__.py +294 -0
  7. sqlalchemy/connectors/__init__.py +18 -0
  8. sqlalchemy/connectors/aioodbc.py +174 -0
  9. sqlalchemy/connectors/asyncio.py +213 -0
  10. sqlalchemy/connectors/pyodbc.py +249 -0
  11. sqlalchemy/cyextension/__init__.py +6 -0
  12. sqlalchemy/cyextension/collections.cp313-win32.pyd +0 -0
  13. sqlalchemy/cyextension/collections.pyx +409 -0
  14. sqlalchemy/cyextension/immutabledict.cp313-win32.pyd +0 -0
  15. sqlalchemy/cyextension/immutabledict.pxd +8 -0
  16. sqlalchemy/cyextension/immutabledict.pyx +133 -0
  17. sqlalchemy/cyextension/processors.cp313-win32.pyd +0 -0
  18. sqlalchemy/cyextension/processors.pyx +68 -0
  19. sqlalchemy/cyextension/resultproxy.cp313-win32.pyd +0 -0
  20. sqlalchemy/cyextension/resultproxy.pyx +102 -0
  21. sqlalchemy/cyextension/util.cp313-win32.pyd +0 -0
  22. sqlalchemy/cyextension/util.pyx +91 -0
  23. sqlalchemy/dialects/__init__.py +61 -0
  24. sqlalchemy/dialects/_typing.py +25 -0
  25. sqlalchemy/dialects/mssql/__init__.py +88 -0
  26. sqlalchemy/dialects/mssql/aioodbc.py +64 -0
  27. sqlalchemy/dialects/mssql/base.py +4010 -0
  28. sqlalchemy/dialects/mssql/information_schema.py +254 -0
  29. sqlalchemy/dialects/mssql/json.py +133 -0
  30. sqlalchemy/dialects/mssql/provision.py +162 -0
  31. sqlalchemy/dialects/mssql/pymssql.py +126 -0
  32. sqlalchemy/dialects/mssql/pyodbc.py +745 -0
  33. sqlalchemy/dialects/mysql/__init__.py +101 -0
  34. sqlalchemy/dialects/mysql/aiomysql.py +333 -0
  35. sqlalchemy/dialects/mysql/asyncmy.py +337 -0
  36. sqlalchemy/dialects/mysql/base.py +3494 -0
  37. sqlalchemy/dialects/mysql/cymysql.py +84 -0
  38. sqlalchemy/dialects/mysql/dml.py +219 -0
  39. sqlalchemy/dialects/mysql/enumerated.py +244 -0
  40. sqlalchemy/dialects/mysql/expression.py +141 -0
  41. sqlalchemy/dialects/mysql/json.py +81 -0
  42. sqlalchemy/dialects/mysql/mariadb.py +32 -0
  43. sqlalchemy/dialects/mysql/mariadbconnector.py +277 -0
  44. sqlalchemy/dialects/mysql/mysqlconnector.py +180 -0
  45. sqlalchemy/dialects/mysql/mysqldb.py +303 -0
  46. sqlalchemy/dialects/mysql/provision.py +110 -0
  47. sqlalchemy/dialects/mysql/pymysql.py +137 -0
  48. sqlalchemy/dialects/mysql/pyodbc.py +138 -0
  49. sqlalchemy/dialects/mysql/reflection.py +677 -0
  50. sqlalchemy/dialects/mysql/reserved_words.py +571 -0
  51. sqlalchemy/dialects/mysql/types.py +774 -0
  52. sqlalchemy/dialects/oracle/__init__.py +67 -0
  53. sqlalchemy/dialects/oracle/base.py +3271 -0
  54. sqlalchemy/dialects/oracle/cx_oracle.py +1483 -0
  55. sqlalchemy/dialects/oracle/dictionary.py +507 -0
  56. sqlalchemy/dialects/oracle/oracledb.py +431 -0
  57. sqlalchemy/dialects/oracle/provision.py +220 -0
  58. sqlalchemy/dialects/oracle/types.py +287 -0
  59. sqlalchemy/dialects/postgresql/__init__.py +167 -0
  60. sqlalchemy/dialects/postgresql/_psycopg_common.py +187 -0
  61. sqlalchemy/dialects/postgresql/array.py +425 -0
  62. sqlalchemy/dialects/postgresql/asyncpg.py +1274 -0
  63. sqlalchemy/dialects/postgresql/base.py +5008 -0
  64. sqlalchemy/dialects/postgresql/dml.py +310 -0
  65. sqlalchemy/dialects/postgresql/ext.py +496 -0
  66. sqlalchemy/dialects/postgresql/hstore.py +397 -0
  67. sqlalchemy/dialects/postgresql/json.py +333 -0
  68. sqlalchemy/dialects/postgresql/named_types.py +509 -0
  69. sqlalchemy/dialects/postgresql/operators.py +129 -0
  70. sqlalchemy/dialects/postgresql/pg8000.py +662 -0
  71. sqlalchemy/dialects/postgresql/pg_catalog.py +300 -0
  72. sqlalchemy/dialects/postgresql/provision.py +175 -0
  73. sqlalchemy/dialects/postgresql/psycopg.py +772 -0
  74. sqlalchemy/dialects/postgresql/psycopg2.py +886 -0
  75. sqlalchemy/dialects/postgresql/psycopg2cffi.py +61 -0
  76. sqlalchemy/dialects/postgresql/ranges.py +1029 -0
  77. sqlalchemy/dialects/postgresql/types.py +303 -0
  78. sqlalchemy/dialects/sqlite/__init__.py +57 -0
  79. sqlalchemy/dialects/sqlite/aiosqlite.py +396 -0
  80. sqlalchemy/dialects/sqlite/base.py +2805 -0
  81. sqlalchemy/dialects/sqlite/dml.py +240 -0
  82. sqlalchemy/dialects/sqlite/json.py +92 -0
  83. sqlalchemy/dialects/sqlite/provision.py +198 -0
  84. sqlalchemy/dialects/sqlite/pysqlcipher.py +155 -0
  85. sqlalchemy/dialects/sqlite/pysqlite.py +756 -0
  86. sqlalchemy/dialects/type_migration_guidelines.txt +145 -0
  87. sqlalchemy/engine/__init__.py +62 -0
  88. sqlalchemy/engine/_py_processors.py +136 -0
  89. sqlalchemy/engine/_py_row.py +128 -0
  90. sqlalchemy/engine/_py_util.py +74 -0
  91. sqlalchemy/engine/base.py +3375 -0
  92. sqlalchemy/engine/characteristics.py +155 -0
  93. sqlalchemy/engine/create.py +875 -0
  94. sqlalchemy/engine/cursor.py +2181 -0
  95. sqlalchemy/engine/default.py +2365 -0
  96. sqlalchemy/engine/events.py +951 -0
  97. sqlalchemy/engine/interfaces.py +3403 -0
  98. sqlalchemy/engine/mock.py +131 -0
  99. sqlalchemy/engine/processors.py +61 -0
  100. sqlalchemy/engine/reflection.py +2098 -0
  101. sqlalchemy/engine/result.py +2382 -0
  102. sqlalchemy/engine/row.py +401 -0
  103. sqlalchemy/engine/strategies.py +19 -0
  104. sqlalchemy/engine/url.py +910 -0
  105. sqlalchemy/engine/util.py +167 -0
  106. sqlalchemy/event/__init__.py +25 -0
  107. sqlalchemy/event/api.py +225 -0
  108. sqlalchemy/event/attr.py +655 -0
  109. sqlalchemy/event/base.py +470 -0
  110. sqlalchemy/event/legacy.py +246 -0
  111. sqlalchemy/event/registry.py +386 -0
  112. sqlalchemy/events.py +17 -0
  113. sqlalchemy/exc.py +830 -0
  114. sqlalchemy/ext/__init__.py +11 -0
  115. sqlalchemy/ext/associationproxy.py +2013 -0
  116. sqlalchemy/ext/asyncio/__init__.py +25 -0
  117. sqlalchemy/ext/asyncio/base.py +279 -0
  118. sqlalchemy/ext/asyncio/engine.py +1466 -0
  119. sqlalchemy/ext/asyncio/exc.py +21 -0
  120. sqlalchemy/ext/asyncio/result.py +961 -0
  121. sqlalchemy/ext/asyncio/scoping.py +1614 -0
  122. sqlalchemy/ext/asyncio/session.py +1936 -0
  123. sqlalchemy/ext/automap.py +1691 -0
  124. sqlalchemy/ext/baked.py +574 -0
  125. sqlalchemy/ext/compiler.py +570 -0
  126. sqlalchemy/ext/declarative/__init__.py +65 -0
  127. sqlalchemy/ext/declarative/extensions.py +548 -0
  128. sqlalchemy/ext/horizontal_shard.py +481 -0
  129. sqlalchemy/ext/hybrid.py +1514 -0
  130. sqlalchemy/ext/indexable.py +341 -0
  131. sqlalchemy/ext/instrumentation.py +450 -0
  132. sqlalchemy/ext/mutable.py +1073 -0
  133. sqlalchemy/ext/mypy/__init__.py +6 -0
  134. sqlalchemy/ext/mypy/apply.py +320 -0
  135. sqlalchemy/ext/mypy/decl_class.py +515 -0
  136. sqlalchemy/ext/mypy/infer.py +590 -0
  137. sqlalchemy/ext/mypy/names.py +335 -0
  138. sqlalchemy/ext/mypy/plugin.py +303 -0
  139. sqlalchemy/ext/mypy/util.py +357 -0
  140. sqlalchemy/ext/orderinglist.py +416 -0
  141. sqlalchemy/ext/serializer.py +181 -0
  142. sqlalchemy/future/__init__.py +16 -0
  143. sqlalchemy/future/engine.py +15 -0
  144. sqlalchemy/inspection.py +174 -0
  145. sqlalchemy/log.py +288 -0
  146. sqlalchemy/orm/__init__.py +170 -0
  147. sqlalchemy/orm/_orm_constructors.py +2571 -0
  148. sqlalchemy/orm/_typing.py +179 -0
  149. sqlalchemy/orm/attributes.py +2835 -0
  150. sqlalchemy/orm/base.py +973 -0
  151. sqlalchemy/orm/bulk_persistence.py +2123 -0
  152. sqlalchemy/orm/clsregistry.py +571 -0
  153. sqlalchemy/orm/collections.py +1620 -0
  154. sqlalchemy/orm/context.py +3268 -0
  155. sqlalchemy/orm/decl_api.py +1883 -0
  156. sqlalchemy/orm/decl_base.py +2190 -0
  157. sqlalchemy/orm/dependency.py +1304 -0
  158. sqlalchemy/orm/descriptor_props.py +1076 -0
  159. sqlalchemy/orm/dynamic.py +300 -0
  160. sqlalchemy/orm/evaluator.py +379 -0
  161. sqlalchemy/orm/events.py +3261 -0
  162. sqlalchemy/orm/exc.py +228 -0
  163. sqlalchemy/orm/identity.py +302 -0
  164. sqlalchemy/orm/instrumentation.py +754 -0
  165. sqlalchemy/orm/interfaces.py +1474 -0
  166. sqlalchemy/orm/loading.py +1682 -0
  167. sqlalchemy/orm/mapped_collection.py +557 -0
  168. sqlalchemy/orm/mapper.py +4432 -0
  169. sqlalchemy/orm/path_registry.py +811 -0
  170. sqlalchemy/orm/persistence.py +1782 -0
  171. sqlalchemy/orm/properties.py +886 -0
  172. sqlalchemy/orm/query.py +3396 -0
  173. sqlalchemy/orm/relationships.py +3500 -0
  174. sqlalchemy/orm/scoping.py +2165 -0
  175. sqlalchemy/orm/session.py +5301 -0
  176. sqlalchemy/orm/state.py +1143 -0
  177. sqlalchemy/orm/state_changes.py +198 -0
  178. sqlalchemy/orm/strategies.py +3473 -0
  179. sqlalchemy/orm/strategy_options.py +2569 -0
  180. sqlalchemy/orm/sync.py +164 -0
  181. sqlalchemy/orm/unitofwork.py +796 -0
  182. sqlalchemy/orm/util.py +2424 -0
  183. sqlalchemy/orm/writeonly.py +678 -0
  184. sqlalchemy/pool/__init__.py +44 -0
  185. sqlalchemy/pool/base.py +1515 -0
  186. sqlalchemy/pool/events.py +370 -0
  187. sqlalchemy/pool/impl.py +581 -0
  188. sqlalchemy/py.typed +0 -0
  189. sqlalchemy/schema.py +70 -0
  190. sqlalchemy/sql/__init__.py +145 -0
  191. sqlalchemy/sql/_dml_constructors.py +140 -0
  192. sqlalchemy/sql/_elements_constructors.py +1850 -0
  193. sqlalchemy/sql/_orm_types.py +20 -0
  194. sqlalchemy/sql/_py_util.py +75 -0
  195. sqlalchemy/sql/_selectable_constructors.py +635 -0
  196. sqlalchemy/sql/_typing.py +460 -0
  197. sqlalchemy/sql/annotation.py +585 -0
  198. sqlalchemy/sql/base.py +2185 -0
  199. sqlalchemy/sql/cache_key.py +1057 -0
  200. sqlalchemy/sql/coercions.py +1405 -0
  201. sqlalchemy/sql/compiler.py +7818 -0
  202. sqlalchemy/sql/crud.py +1669 -0
  203. sqlalchemy/sql/ddl.py +1378 -0
  204. sqlalchemy/sql/default_comparator.py +552 -0
  205. sqlalchemy/sql/dml.py +1817 -0
  206. sqlalchemy/sql/elements.py +5499 -0
  207. sqlalchemy/sql/events.py +455 -0
  208. sqlalchemy/sql/expression.py +162 -0
  209. sqlalchemy/sql/functions.py +2055 -0
  210. sqlalchemy/sql/lambdas.py +1449 -0
  211. sqlalchemy/sql/naming.py +212 -0
  212. sqlalchemy/sql/operators.py +2579 -0
  213. sqlalchemy/sql/roles.py +323 -0
  214. sqlalchemy/sql/schema.py +6158 -0
  215. sqlalchemy/sql/selectable.py +7004 -0
  216. sqlalchemy/sql/sqltypes.py +3827 -0
  217. sqlalchemy/sql/traversals.py +1024 -0
  218. sqlalchemy/sql/type_api.py +2339 -0
  219. sqlalchemy/sql/util.py +1486 -0
  220. sqlalchemy/sql/visitors.py +1165 -0
  221. sqlalchemy/testing/__init__.py +96 -0
  222. sqlalchemy/testing/assertions.py +989 -0
  223. sqlalchemy/testing/assertsql.py +516 -0
  224. sqlalchemy/testing/asyncio.py +135 -0
  225. sqlalchemy/testing/config.py +427 -0
  226. sqlalchemy/testing/engines.py +472 -0
  227. sqlalchemy/testing/entities.py +117 -0
  228. sqlalchemy/testing/exclusions.py +435 -0
  229. sqlalchemy/testing/fixtures/__init__.py +28 -0
  230. sqlalchemy/testing/fixtures/base.py +366 -0
  231. sqlalchemy/testing/fixtures/mypy.py +312 -0
  232. sqlalchemy/testing/fixtures/orm.py +227 -0
  233. sqlalchemy/testing/fixtures/sql.py +503 -0
  234. sqlalchemy/testing/pickleable.py +155 -0
  235. sqlalchemy/testing/plugin/__init__.py +6 -0
  236. sqlalchemy/testing/plugin/bootstrap.py +51 -0
  237. sqlalchemy/testing/plugin/plugin_base.py +779 -0
  238. sqlalchemy/testing/plugin/pytestplugin.py +868 -0
  239. sqlalchemy/testing/profiling.py +324 -0
  240. sqlalchemy/testing/provision.py +496 -0
  241. sqlalchemy/testing/requirements.py +1818 -0
  242. sqlalchemy/testing/schema.py +224 -0
  243. sqlalchemy/testing/suite/__init__.py +19 -0
  244. sqlalchemy/testing/suite/test_cte.py +211 -0
  245. sqlalchemy/testing/suite/test_ddl.py +389 -0
  246. sqlalchemy/testing/suite/test_deprecations.py +153 -0
  247. sqlalchemy/testing/suite/test_dialect.py +740 -0
  248. sqlalchemy/testing/suite/test_insert.py +630 -0
  249. sqlalchemy/testing/suite/test_reflection.py +3225 -0
  250. sqlalchemy/testing/suite/test_results.py +502 -0
  251. sqlalchemy/testing/suite/test_rowcount.py +258 -0
  252. sqlalchemy/testing/suite/test_select.py +1999 -0
  253. sqlalchemy/testing/suite/test_sequence.py +317 -0
  254. sqlalchemy/testing/suite/test_types.py +2141 -0
  255. sqlalchemy/testing/suite/test_unicode_ddl.py +189 -0
  256. sqlalchemy/testing/suite/test_update_delete.py +139 -0
  257. sqlalchemy/testing/util.py +537 -0
  258. sqlalchemy/testing/warnings.py +52 -0
  259. sqlalchemy/types.py +76 -0
  260. sqlalchemy/util/__init__.py +160 -0
  261. sqlalchemy/util/_collections.py +715 -0
  262. sqlalchemy/util/_concurrency_py3k.py +288 -0
  263. sqlalchemy/util/_has_cy.py +40 -0
  264. sqlalchemy/util/_py_collections.py +541 -0
  265. sqlalchemy/util/compat.py +301 -0
  266. sqlalchemy/util/concurrency.py +108 -0
  267. sqlalchemy/util/deprecations.py +401 -0
  268. sqlalchemy/util/langhelpers.py +2218 -0
  269. sqlalchemy/util/preloaded.py +150 -0
  270. sqlalchemy/util/queue.py +322 -0
  271. sqlalchemy/util/tool_support.py +201 -0
  272. sqlalchemy/util/topological.py +120 -0
  273. sqlalchemy/util/typing.py +629 -0
@@ -0,0 +1,516 @@
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
+
10
+ from __future__ import annotations
11
+
12
+ import collections
13
+ import contextlib
14
+ import itertools
15
+ import re
16
+
17
+ from .. import event
18
+ from ..engine import url
19
+ from ..engine.default import DefaultDialect
20
+ from ..schema import BaseDDLElement
21
+
22
+
23
+ class 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
+
38
+ class SQLMatchRule(AssertRule):
39
+ pass
40
+
41
+
42
+ class 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
+
70
+ class 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
+
237
+ class 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
+
263
+ class 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
+
339
+ class 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
+
355
+ class 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
+
376
+ class 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
+
405
+ class 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
+
413
+ class 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
+
424
+ class 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
+
441
+ class SQLCursorExecuteObserved(
442
+ collections.namedtuple(
443
+ "SQLCursorExecuteObserved",
444
+ ["statement", "parameters", "context", "executemany"],
445
+ )
446
+ ):
447
+ pass
448
+
449
+
450
+ class 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
476
+ def 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()
@@ -0,0 +1,135 @@
1
+ # testing/asyncio.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
+
10
+ # functions and wrappers to run tests, fixtures, provisioning and
11
+ # setup/teardown in an asyncio event loop, conditionally based on the
12
+ # current DB driver being used for a test.
13
+
14
+ # note that SQLAlchemy's asyncio integration also supports a method
15
+ # of running individual asyncio functions inside of separate event loops
16
+ # using "async_fallback" mode; however running whole functions in the event
17
+ # loop is a more accurate test for how SQLAlchemy's asyncio features
18
+ # would run in the real world.
19
+
20
+
21
+ from __future__ import annotations
22
+
23
+ from functools import wraps
24
+ import inspect
25
+
26
+ from . import config
27
+ from ..util.concurrency import _AsyncUtil
28
+
29
+ # may be set to False if the
30
+ # --disable-asyncio flag is passed to the test runner.
31
+ ENABLE_ASYNCIO = True
32
+ _async_util = _AsyncUtil() # it has lazy init so just always create one
33
+
34
+
35
+ def _shutdown():
36
+ """called when the test finishes"""
37
+ _async_util.close()
38
+
39
+
40
+ def _run_coroutine_function(fn, *args, **kwargs):
41
+ return _async_util.run(fn, *args, **kwargs)
42
+
43
+
44
+ def _assume_async(fn, *args, **kwargs):
45
+ """Run a function in an asyncio loop unconditionally.
46
+
47
+ This function is used for provisioning features like
48
+ testing a database connection for server info.
49
+
50
+ Note that for blocking IO database drivers, this means they block the
51
+ event loop.
52
+
53
+ """
54
+
55
+ if not ENABLE_ASYNCIO:
56
+ return fn(*args, **kwargs)
57
+
58
+ return _async_util.run_in_greenlet(fn, *args, **kwargs)
59
+
60
+
61
+ def _maybe_async_provisioning(fn, *args, **kwargs):
62
+ """Run a function in an asyncio loop if any current drivers might need it.
63
+
64
+ This function is used for provisioning features that take
65
+ place outside of a specific database driver being selected, so if the
66
+ current driver that happens to be used for the provisioning operation
67
+ is an async driver, it will run in asyncio and not fail.
68
+
69
+ Note that for blocking IO database drivers, this means they block the
70
+ event loop.
71
+
72
+ """
73
+ if not ENABLE_ASYNCIO:
74
+ return fn(*args, **kwargs)
75
+
76
+ if config.any_async:
77
+ return _async_util.run_in_greenlet(fn, *args, **kwargs)
78
+ else:
79
+ return fn(*args, **kwargs)
80
+
81
+
82
+ def _maybe_async(fn, *args, **kwargs):
83
+ """Run a function in an asyncio loop if the current selected driver is
84
+ async.
85
+
86
+ This function is used for test setup/teardown and tests themselves
87
+ where the current DB driver is known.
88
+
89
+
90
+ """
91
+ if not ENABLE_ASYNCIO:
92
+ return fn(*args, **kwargs)
93
+
94
+ is_async = config._current.is_async
95
+
96
+ if is_async:
97
+ return _async_util.run_in_greenlet(fn, *args, **kwargs)
98
+ else:
99
+ return fn(*args, **kwargs)
100
+
101
+
102
+ def _maybe_async_wrapper(fn):
103
+ """Apply the _maybe_async function to an existing function and return
104
+ as a wrapped callable, supporting generator functions as well.
105
+
106
+ This is currently used for pytest fixtures that support generator use.
107
+
108
+ """
109
+
110
+ if inspect.isgeneratorfunction(fn):
111
+ _stop = object()
112
+
113
+ def call_next(gen):
114
+ try:
115
+ return next(gen)
116
+ # can't raise StopIteration in an awaitable.
117
+ except StopIteration:
118
+ return _stop
119
+
120
+ @wraps(fn)
121
+ def wrap_fixture(*args, **kwargs):
122
+ gen = fn(*args, **kwargs)
123
+ while True:
124
+ value = _maybe_async(call_next, gen)
125
+ if value is _stop:
126
+ break
127
+ yield value
128
+
129
+ else:
130
+
131
+ @wraps(fn)
132
+ def wrap_fixture(*args, **kwargs):
133
+ return _maybe_async(fn, *args, **kwargs)
134
+
135
+ return wrap_fixture