codekingpro/portable-devtools
114k
1# testing/suite/test_dialect.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
10import importlib
11
12from . import testing
13from .. import assert_raises
14from .. import config
15from .. import engines
16from .. import eq_
17from .. import fixtures
18from .. import is_not_none
19from .. import is_true
20from .. import mock
21from .. import ne_
22from .. import provide_metadata
23from ..assertions import expect_raises
24from ..assertions import expect_raises_message
25from ..config import requirements
26from ..provision import set_default_schema_on_connection
27from ..schema import Column
28from ..schema import Table
29from ... import bindparam
30from ... import dialects
31from ... import event
32from ... import exc
33from ... import Integer
34from ... import literal_column
35from ... import select
36from ... import String
37from ...sql.compiler import Compiled
38from ...util import inspect_getfullargspec
39
40
41class PingTest(fixtures.TestBase):
42 __backend__ = True
43
44 def test_do_ping(self):
45 with testing.db.connect() as conn:
46 is_true(
47 testing.db.dialect.do_ping(conn.connection.dbapi_connection)
48 )
49
50
51class ArgSignatureTest(fixtures.TestBase):
52 """test that all visit_XYZ() in :class:`_sql.Compiler` subclasses have
53 ``**kw``, for #8988.
54
55 This test uses runtime code inspection. Does not need to be a
56 ``__backend__`` test as it only needs to run once provided all target
57 dialects have been imported.
58
59 For third party dialects, the suite would be run with that third
60 party as a "--dburi", which means its compiler classes will have been
61 imported by the time this test runs.
62
63 """
64
65 def _all_subclasses(): # type: ignore # noqa
66 for d in dialects.__all__:
67 if not d.startswith("_"):
68 importlib.import_module("sqlalchemy.dialects.%s" % d)
69
70 stack = [Compiled]
71
72 while stack:
73 cls = stack.pop(0)
74 stack.extend(cls.__subclasses__())
75 yield cls
76
77 @testing.fixture(params=list(_all_subclasses()))
78 def all_subclasses(self, request):
79 yield request.param
80
81 def test_all_visit_methods_accept_kw(self, all_subclasses):
82 cls = all_subclasses
83
84 for k in cls.__dict__:
85 if k.startswith("visit_"):
86 meth = getattr(cls, k)
87
88 insp = inspect_getfullargspec(meth)
89 is_not_none(
90 insp.varkw,
91 f"Compiler visit method {cls.__name__}.{k}() does "
92 "not accommodate for **kw in its argument signature",
93 )
94
95
96class ExceptionTest(fixtures.TablesTest):
97 """Test basic exception wrapping.
98
99 DBAPIs vary a lot in exception behavior so to actually anticipate
100 specific exceptions from real round trips, we need to be conservative.
101
102 """
103
104 run_deletes = "each"
105
106 __backend__ = True
107
108 @classmethod
109 def define_tables(cls, metadata):
110 Table(
111 "manual_pk",
112 metadata,
113 Column("id", Integer, primary_key=True, autoincrement=False),
114 Column("data", String(50)),
115 )
116
117 @requirements.duplicate_key_raises_integrity_error
118 def test_integrity_error(self):
119 with config.db.connect() as conn:
120 trans = conn.begin()
121 conn.execute(
122 self.tables.manual_pk.insert(), {"id": 1, "data": "d1"}
123 )
124
125 assert_raises(
126 exc.IntegrityError,
127 conn.execute,
128 self.tables.manual_pk.insert(),
129 {"id": 1, "data": "d1"},
130 )
131
132 trans.rollback()
133
134 def test_exception_with_non_ascii(self):
135 with config.db.connect() as conn:
136 try:
137 # try to create an error message that likely has non-ascii
138 # characters in the DBAPI's message string. unfortunately
139 # there's no way to make this happen with some drivers like
140 # mysqlclient, pymysql. this at least does produce a non-
141 # ascii error message for cx_oracle, psycopg2
142 conn.execute(select(literal_column("méil")))
143 assert False
144 except exc.DBAPIError as err:
145 err_str = str(err)
146
147 assert str(err.orig) in str(err)
148
149 assert isinstance(err_str, str)
150
151
152class IsolationLevelTest(fixtures.TestBase):
153 __backend__ = True
154
155 __requires__ = ("isolation_level",)
156
157 def _get_non_default_isolation_level(self):
158 levels = requirements.get_isolation_levels(config)
159
160 default = levels["default"]
161 supported = levels["supported"]
162
163 s = set(supported).difference(["AUTOCOMMIT", default])
164 if s:
165 return s.pop()
166 else:
167 config.skip_test("no non-default isolation level available")
168
169 def test_default_isolation_level(self):
170 eq_(
171 config.db.dialect.default_isolation_level,
172 requirements.get_isolation_levels(config)["default"],
173 )
174
175 def test_non_default_isolation_level(self):
176 non_default = self._get_non_default_isolation_level()
177
178 with config.db.connect() as conn:
179 existing = conn.get_isolation_level()
180
181 ne_(existing, non_default)
182
183 conn.execution_options(isolation_level=non_default)
184
185 eq_(conn.get_isolation_level(), non_default)
186
187 conn.dialect.reset_isolation_level(
188 conn.connection.dbapi_connection
189 )
190
191 eq_(conn.get_isolation_level(), existing)
192
193 def test_all_levels(self):
194 levels = requirements.get_isolation_levels(config)
195
196 all_levels = levels["supported"]
197
198 for level in set(all_levels).difference(["AUTOCOMMIT"]):
199 with config.db.connect() as conn:
200 conn.execution_options(isolation_level=level)
201
202 eq_(conn.get_isolation_level(), level)
203
204 trans = conn.begin()
205 trans.rollback()
206
207 eq_(conn.get_isolation_level(), level)
208
209 with config.db.connect() as conn:
210 eq_(
211 conn.get_isolation_level(),
212 levels["default"],
213 )
214
215 @testing.requires.get_isolation_level_values
216 def test_invalid_level_execution_option(self, connection_no_trans):
217 """test for the new get_isolation_level_values() method"""
218
219 connection = connection_no_trans
220 with expect_raises_message(
221 exc.ArgumentError,
222 "Invalid value '%s' for isolation_level. "
223 "Valid isolation levels for '%s' are %s"
224 % (
225 "FOO",
226 connection.dialect.name,
227 ", ".join(
228 requirements.get_isolation_levels(config)["supported"]
229 ),
230 ),
231 ):
232 connection.execution_options(isolation_level="FOO")
233
234 @testing.requires.get_isolation_level_values
235 @testing.requires.dialect_level_isolation_level_param
236 def test_invalid_level_engine_param(self, testing_engine):
237 """test for the new get_isolation_level_values() method
238 and support for the dialect-level 'isolation_level' parameter.
239
240 """
241
242 eng = testing_engine(options=dict(isolation_level="FOO"))
243 with expect_raises_message(
244 exc.ArgumentError,
245 "Invalid value '%s' for isolation_level. "
246 "Valid isolation levels for '%s' are %s"
247 % (
248 "FOO",
249 eng.dialect.name,
250 ", ".join(
251 requirements.get_isolation_levels(config)["supported"]
252 ),
253 ),
254 ):
255 eng.connect()
256
257 @testing.requires.independent_readonly_connections
258 def test_dialect_user_setting_is_restored(self, testing_engine):
259 levels = requirements.get_isolation_levels(config)
260 default = levels["default"]
261 supported = (
262 sorted(
263 set(levels["supported"]).difference([default, "AUTOCOMMIT"])
264 )
265 )[0]
266
267 e = testing_engine(options={"isolation_level": supported})
268
269 with e.connect() as conn:
270 eq_(conn.get_isolation_level(), supported)
271
272 with e.connect() as conn:
273 conn.execution_options(isolation_level=default)
274 eq_(conn.get_isolation_level(), default)
275
276 with e.connect() as conn:
277 eq_(conn.get_isolation_level(), supported)
278
279
280class AutocommitIsolationTest(fixtures.TablesTest):
281 run_deletes = "each"
282
283 __requires__ = ("autocommit",)
284
285 __backend__ = True
286
287 @classmethod
288 def define_tables(cls, metadata):
289 Table(
290 "some_table",
291 metadata,
292 Column("id", Integer, primary_key=True, autoincrement=False),
293 Column("data", String(50)),
294 test_needs_acid=True,
295 )
296
297 def _test_conn_autocommits(self, conn, autocommit, ensure_table=False):
298 if ensure_table:
299 self.tables.some_table.create(conn, checkfirst=True)
300 conn.commit()
301
302 trans = conn.begin()
303 conn.execute(
304 self.tables.some_table.insert(), {"id": 1, "data": "some data"}
305 )
306 trans.rollback()
307
308 eq_(
309 conn.scalar(select(self.tables.some_table.c.id)),
310 1 if autocommit else None,
311 )
312 conn.rollback()
313
314 with conn.begin():
315 conn.execute(self.tables.some_table.delete())
316
317 def test_autocommit_on(self, connection_no_trans):
318 conn = connection_no_trans
319 c2 = conn.execution_options(isolation_level="AUTOCOMMIT")
320 self._test_conn_autocommits(c2, True)
321
322 c2.dialect.reset_isolation_level(c2.connection.dbapi_connection)
323
324 self._test_conn_autocommits(conn, False)
325
326 def test_autocommit_off(self, connection_no_trans):
327 conn = connection_no_trans
328 self._test_conn_autocommits(conn, False)
329
330 def test_turn_autocommit_off_via_default_iso_level(
331 self, connection_no_trans
332 ):
333 conn = connection_no_trans
334 conn = conn.execution_options(isolation_level="AUTOCOMMIT")
335 self._test_conn_autocommits(conn, True)
336
337 conn.execution_options(
338 isolation_level=requirements.get_isolation_levels(config)[
339 "default"
340 ]
341 )
342 self._test_conn_autocommits(conn, False)
343
344 @testing.requires.skip_autocommit_rollback
345 @testing.variation("autocommit_setting", ["false", "engine", "option"])
346 @testing.variation("block_rollback", [True, False])
347 def test_autocommit_block(
348 self, testing_engine, autocommit_setting, block_rollback
349 ):
350 kw = {}
351 if bool(block_rollback):
352 kw["skip_autocommit_rollback"] = True
353 if autocommit_setting.engine:
354 kw["isolation_level"] = "AUTOCOMMIT"
355
356 engine = testing_engine(options=kw)
357
358 conn = engine.connect()
359 if autocommit_setting.option:
360 conn.execution_options(isolation_level="AUTOCOMMIT")
361 self._test_conn_autocommits(
362 conn,
363 autocommit_setting.engine or autocommit_setting.option,
364 ensure_table=True,
365 )
366 with mock.patch.object(
367 conn.connection, "rollback", wraps=conn.connection.rollback
368 ) as check_rollback:
369 conn.close()
370 if autocommit_setting.false or not block_rollback:
371 eq_(check_rollback.mock_calls, [mock.call()])
372 else:
373 eq_(check_rollback.mock_calls, [])
374
375 @testing.requires.independent_readonly_connections
376 @testing.variation("use_dialect_setting", [True, False])
377 def test_dialect_autocommit_is_restored(
378 self, testing_engine, use_dialect_setting
379 ):
380 """test #10147"""
381
382 if use_dialect_setting:
383 e = testing_engine(options={"isolation_level": "AUTOCOMMIT"})
384 else:
385 e = testing_engine().execution_options(
386 isolation_level="AUTOCOMMIT"
387 )
388
389 levels = requirements.get_isolation_levels(config)
390
391 default = levels["default"]
392
393 with e.connect() as conn:
394 self._test_conn_autocommits(conn, True)
395
396 with e.connect() as conn:
397 conn.execution_options(isolation_level=default)
398 self._test_conn_autocommits(conn, False)
399
400 with e.connect() as conn:
401 self._test_conn_autocommits(conn, True)
402
403
404class EscapingTest(fixtures.TestBase):
405 @provide_metadata
406 def test_percent_sign_round_trip(self):
407 """test that the DBAPI accommodates for escaped / nonescaped
408 percent signs in a way that matches the compiler
409
410 """
411 m = self.metadata
412 t = Table("t", m, Column("data", String(50)))
413 t.create(config.db)
414 with config.db.begin() as conn:
415 conn.execute(t.insert(), dict(data="some % value"))
416 conn.execute(t.insert(), dict(data="some %% other value"))
417
418 eq_(
419 conn.scalar(
420 select(t.c.data).where(
421 t.c.data == literal_column("'some % value'")
422 )
423 ),
424 "some % value",
425 )
426
427 eq_(
428 conn.scalar(
429 select(t.c.data).where(
430 t.c.data == literal_column("'some %% other value'")
431 )
432 ),
433 "some %% other value",
434 )
435
436
437class WeCanSetDefaultSchemaWEventsTest(fixtures.TestBase):
438 __backend__ = True
439
440 __requires__ = ("default_schema_name_switch",)
441
442 def test_control_case(self):
443 default_schema_name = config.db.dialect.default_schema_name
444
445 eng = engines.testing_engine()
446 with eng.connect():
447 pass
448
449 eq_(eng.dialect.default_schema_name, default_schema_name)
450
451 def test_wont_work_wo_insert(self):
452 default_schema_name = config.db.dialect.default_schema_name
453
454 eng = engines.testing_engine()
455
456 @event.listens_for(eng, "connect")
457 def on_connect(dbapi_connection, connection_record):
458 set_default_schema_on_connection(
459 config, dbapi_connection, config.test_schema
460 )
461
462 with eng.connect() as conn:
463 what_it_should_be = eng.dialect._get_default_schema_name(conn)
464 eq_(what_it_should_be, config.test_schema)
465
466 eq_(eng.dialect.default_schema_name, default_schema_name)
467
468 def test_schema_change_on_connect(self):
469 eng = engines.testing_engine()
470
471 @event.listens_for(eng, "connect", insert=True)
472 def on_connect(dbapi_connection, connection_record):
473 set_default_schema_on_connection(
474 config, dbapi_connection, config.test_schema
475 )
476
477 with eng.connect() as conn:
478 what_it_should_be = eng.dialect._get_default_schema_name(conn)
479 eq_(what_it_should_be, config.test_schema)
480
481 eq_(eng.dialect.default_schema_name, config.test_schema)
482
483 def test_schema_change_works_w_transactions(self):
484 eng = engines.testing_engine()
485
486 @event.listens_for(eng, "connect", insert=True)
487 def on_connect(dbapi_connection, *arg):
488 set_default_schema_on_connection(
489 config, dbapi_connection, config.test_schema
490 )
491
492 with eng.connect() as conn:
493 trans = conn.begin()
494 what_it_should_be = eng.dialect._get_default_schema_name(conn)
495 eq_(what_it_should_be, config.test_schema)
496 trans.rollback()
497
498 what_it_should_be = eng.dialect._get_default_schema_name(conn)
499 eq_(what_it_should_be, config.test_schema)
500
501 eq_(eng.dialect.default_schema_name, config.test_schema)
502
503
504class FutureWeCanSetDefaultSchemaWEventsTest(
505 fixtures.FutureEngineMixin, WeCanSetDefaultSchemaWEventsTest
506):
507 pass
508
509
510class DifficultParametersTest(fixtures.TestBase):
511 __backend__ = True
512
513 tough_parameters = testing.combinations(
514 ("boring",),
515 ("per cent",),
516 ("per % cent",),
517 ("%percent",),
518 ("par(ens)",),
519 ("percent%(ens)yah",),
520 ("col:ons",),
521 ("_starts_with_underscore",),
522 ("dot.s",),
523 ("more :: %colons%",),
524 ("_name",),
525 ("___name",),
526 ("[BracketsAndCase]",),
527 ("42numbers",),
528 ("percent%signs",),
529 ("has spaces",),
530 ("/slashes/",),
531 ("more/slashes",),
532 ("q?marks",),
533 ("1param",),
534 ("1col:on",),
535 argnames="paramname",
536 )
537
538 @tough_parameters
539 @config.requirements.unusual_column_name_characters
540 def test_round_trip_same_named_column(
541 self, paramname, connection, metadata
542 ):
543 name = paramname
544
545 t = Table(
546 "t",
547 metadata,
548 Column("id", Integer, primary_key=True),
549 Column(name, String(50), nullable=False),
550 )
551
552 # table is created
553 t.create(connection)
554
555 # automatic param generated by insert
556 connection.execute(t.insert().values({"id": 1, name: "some name"}))
557
558 # automatic param generated by criteria, plus selecting the column
559 stmt = select(t.c[name]).where(t.c[name] == "some name")
560
561 eq_(connection.scalar(stmt), "some name")
562
563 # use the name in a param explicitly
564 stmt = select(t.c[name]).where(t.c[name] == bindparam(name))
565
566 row = connection.execute(stmt, {name: "some name"}).first()
567
568 # name works as the key from cursor.description
569 eq_(row._mapping[name], "some name")
570
571 # use expanding IN
572 stmt = select(t.c[name]).where(
573 t.c[name].in_(["some name", "some other_name"])
574 )
575
576 connection.execute(stmt).first()
577
578 @testing.fixture
579 def multirow_fixture(self, metadata, connection):
580 mytable = Table(
581 "mytable",
582 metadata,
583 Column("myid", Integer),
584 Column("name", String(50)),
585 Column("desc", String(50)),
586 )
587
588 mytable.create(connection)
589
590 connection.execute(
591 mytable.insert(),
592 [
593 {"myid": 1, "name": "a", "desc": "a_desc"},
594 {"myid": 2, "name": "b", "desc": "b_desc"},
595 {"myid": 3, "name": "c", "desc": "c_desc"},
596 {"myid": 4, "name": "d", "desc": "d_desc"},
597 ],
598 )
599 yield mytable
600
601 @tough_parameters
602 def test_standalone_bindparam_escape(
603 self, paramname, connection, multirow_fixture
604 ):
605 tbl1 = multirow_fixture
606 stmt = select(tbl1.c.myid).where(
607 tbl1.c.name == bindparam(paramname, value="x")
608 )
609 res = connection.scalar(stmt, {paramname: "c"})
610 eq_(res, 3)
611
612 @tough_parameters
613 def test_standalone_bindparam_escape_expanding(
614 self, paramname, connection, multirow_fixture
615 ):
616 tbl1 = multirow_fixture
617 stmt = (
618 select(tbl1.c.myid)
619 .where(tbl1.c.name.in_(bindparam(paramname, value=["a", "b"])))
620 .order_by(tbl1.c.myid)
621 )
622
623 res = connection.scalars(stmt, {paramname: ["d", "a"]}).all()
624 eq_(res, [1, 4])
625
626
627class ReturningGuardsTest(fixtures.TablesTest):
628 """test that the various 'returning' flags are set appropriately"""
629
630 __backend__ = True
631
632 @classmethod
633 def define_tables(cls, metadata):
634 Table(
635 "t",
636 metadata,
637 Column("id", Integer, primary_key=True, autoincrement=False),
638 Column("data", String(50)),
639 )
640
641 @testing.fixture
642 def run_stmt(self, connection):
643 t = self.tables.t
644
645 def go(stmt, executemany, id_param_name, expect_success):
646 stmt = stmt.returning(t.c.id)
647
648 if executemany:
649 if not expect_success:
650 # for RETURNING executemany(), we raise our own
651 # error as this is independent of general RETURNING
652 # support
653 with expect_raises_message(
654 exc.StatementError,
655 rf"Dialect {connection.dialect.name}\+"
656 f"{connection.dialect.driver} with "
657 f"current server capabilities does not support "
658 f".*RETURNING when executemany is used",
659 ):
660 connection.execute(
661 stmt,
662 [
663 {id_param_name: 1, "data": "d1"},
664 {id_param_name: 2, "data": "d2"},
665 {id_param_name: 3, "data": "d3"},
666 ],
667 )
668 else:
669 result = connection.execute(
670 stmt,
671 [
672 {id_param_name: 1, "data": "d1"},
673 {id_param_name: 2, "data": "d2"},
674 {id_param_name: 3, "data": "d3"},
675 ],
676 )
677 eq_(result.all(), [(1,), (2,), (3,)])
678 else:
679 if not expect_success:
680 # for RETURNING execute(), we pass all the way to the DB
681 # and let it fail
682 with expect_raises(exc.DBAPIError):
683 connection.execute(
684 stmt, {id_param_name: 1, "data": "d1"}
685 )
686 else:
687 result = connection.execute(
688 stmt, {id_param_name: 1, "data": "d1"}
689 )
690 eq_(result.all(), [(1,)])
691
692 return go
693
694 def test_insert_single(self, connection, run_stmt):
695 t = self.tables.t
696
697 stmt = t.insert()
698
699 run_stmt(stmt, False, "id", connection.dialect.insert_returning)
700
701 def test_insert_many(self, connection, run_stmt):
702 t = self.tables.t
703
704 stmt = t.insert()
705
706 run_stmt(
707 stmt, True, "id", connection.dialect.insert_executemany_returning
708 )
709
710 def test_update_single(self, connection, run_stmt):
711 t = self.tables.t
712
713 connection.execute(
714 t.insert(),
715 [
716 {"id": 1, "data": "d1"},
717 {"id": 2, "data": "d2"},
718 {"id": 3, "data": "d3"},
719 ],
720 )
721
722 stmt = t.update().where(t.c.id == bindparam("b_id"))
723
724 run_stmt(stmt, False, "b_id", connection.dialect.update_returning)
725
726 def test_update_many(self, connection, run_stmt):
727 t = self.tables.t
728
729 connection.execute(
730 t.insert(),
731 [
732 {"id": 1, "data": "d1"},
733 {"id": 2, "data": "d2"},
734 {"id": 3, "data": "d3"},
735 ],
736 )
737
738 stmt = t.update().where(t.c.id == bindparam("b_id"))
739
740 run_stmt(
741 stmt, True, "b_id", connection.dialect.update_executemany_returning
742 )
743
744 def test_delete_single(self, connection, run_stmt):
745 t = self.tables.t
746
747 connection.execute(
748 t.insert(),
749 [
750 {"id": 1, "data": "d1"},
751 {"id": 2, "data": "d2"},
752 {"id": 3, "data": "d3"},
753 ],
754 )
755
756 stmt = t.delete().where(t.c.id == bindparam("b_id"))
757
758 run_stmt(stmt, False, "b_id", connection.dialect.delete_returning)
759
760 def test_delete_many(self, connection, run_stmt):
761 t = self.tables.t
762
763 connection.execute(
764 t.insert(),
765 [
766 {"id": 1, "data": "d1"},
767 {"id": 2, "data": "d2"},
768 {"id": 3, "data": "d3"},
769 ],
770 )
771
772 stmt = t.delete().where(t.c.id == bindparam("b_id"))
773
774 run_stmt(
775 stmt, True, "b_id", connection.dialect.delete_executemany_returning
776 )
777 