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