Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_dialect.py777 linesDownload Raw Back to suite
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 
codekingpro/portable-devtools · Team Ai