Team Ai
Datasetpublic

codekingpro/portable-devtools

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