Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
default.py2349 linesDownload Raw Back to engine
1# engine/default.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: allow-untyped-defs, allow-untyped-calls
8
9"""Default implementations of per-dialect sqlalchemy.engine classes.
10
11These are semi-private implementation classes which are only of importance
12to database dialect authors; dialects will usually use the classes here
13as the base class for their own corresponding classes.
14
15"""
16
17from __future__ import annotations
18
19import functools
20import operator
21import random
22import re
23from time import perf_counter
24import typing
25from typing import Any
26from typing import Callable
27from typing import cast
28from typing import Dict
29from typing import List
30from typing import Mapping
31from typing import MutableMapping
32from typing import MutableSequence
33from typing import Optional
34from typing import Sequence
35from typing import Set
36from typing import Tuple
37from typing import Type
38from typing import TYPE_CHECKING
39from typing import Union
40import weakref
41
42from . import characteristics
43from . import cursor as _cursor
44from . import interfaces
45from .base import Connection
46from .interfaces import CacheStats
47from .interfaces import DBAPICursor
48from .interfaces import Dialect
49from .interfaces import ExecuteStyle
50from .interfaces import ExecutionContext
51from .reflection import ObjectKind
52from .reflection import ObjectScope
53from .. import event
54from .. import exc
55from .. import pool
56from .. import util
57from ..sql import compiler
58from ..sql import dml
59from ..sql import expression
60from ..sql import type_api
61from ..sql._typing import is_tuple_type
62from ..sql.base import _NoArg
63from ..sql.compiler import DDLCompiler
64from ..sql.compiler import InsertmanyvaluesSentinelOpts
65from ..sql.compiler import SQLCompiler
66from ..sql.elements import quoted_name
67from ..util.typing import Final
68from ..util.typing import Literal
69
70if typing.TYPE_CHECKING:
71    from types import ModuleType
72
73    from .base import Engine
74    from .cursor import ResultFetchStrategy
75    from .interfaces import _CoreMultiExecuteParams
76    from .interfaces import _CoreSingleExecuteParams
77    from .interfaces import _DBAPICursorDescription
78    from .interfaces import _DBAPIMultiExecuteParams
79    from .interfaces import _ExecuteOptions
80    from .interfaces import _MutableCoreSingleExecuteParams
81    from .interfaces import _ParamStyle
82    from .interfaces import DBAPIConnection
83    from .interfaces import IsolationLevel
84    from .row import Row
85    from .url import URL
86    from ..event import _ListenerFnType
87    from ..pool import Pool
88    from ..pool import PoolProxiedConnection
89    from ..sql import Executable
90    from ..sql.compiler import Compiled
91    from ..sql.compiler import Linting
92    from ..sql.compiler import ResultColumnsEntry
93    from ..sql.dml import DMLState
94    from ..sql.dml import UpdateBase
95    from ..sql.elements import BindParameter
96    from ..sql.schema import Column
97    from ..sql.type_api import _BindProcessorType
98    from ..sql.type_api import _ResultProcessorType
99    from ..sql.type_api import TypeEngine
100
101# When we're handed literal SQL, ensure it's a SELECT query
102SERVER_SIDE_CURSOR_RE = re.compile(r"\s*SELECT", re.I | re.UNICODE)
103
104
105(
106    CACHE_HIT,
107    CACHE_MISS,
108    CACHING_DISABLED,
109    NO_CACHE_KEY,
110    NO_DIALECT_SUPPORT,
111) = list(CacheStats)
112
113
114class DefaultDialect(Dialect):
115    """Default implementation of Dialect"""
116
117    statement_compiler = compiler.SQLCompiler
118    ddl_compiler = compiler.DDLCompiler
119    type_compiler_cls = compiler.GenericTypeCompiler
120
121    preparer = compiler.IdentifierPreparer
122    supports_alter = True
123    supports_comments = False
124    supports_constraint_comments = False
125    inline_comments = False
126    supports_statement_cache = True
127
128    div_is_floordiv = True
129
130    bind_typing = interfaces.BindTyping.NONE
131
132    include_set_input_sizes: Optional[Set[Any]] = None
133    exclude_set_input_sizes: Optional[Set[Any]] = None
134
135    # the first value we'd get for an autoincrement column.
136    default_sequence_base = 1
137
138    # most DBAPIs happy with this for execute().
139    # not cx_oracle.
140    execute_sequence_format = tuple
141
142    supports_schemas = True
143    supports_views = True
144    supports_sequences = False
145    sequences_optional = False
146    preexecute_autoincrement_sequences = False
147    supports_identity_columns = False
148    postfetch_lastrowid = True
149    favor_returning_over_lastrowid = False
150    insert_null_pk_still_autoincrements = False
151    update_returning = False
152    delete_returning = False
153    update_returning_multifrom = False
154    delete_returning_multifrom = False
155    insert_returning = False
156
157    cte_follows_insert = False
158
159    supports_native_enum = False
160    supports_native_boolean = False
161    supports_native_uuid = False
162    returns_native_bytes = False
163
164    non_native_boolean_check_constraint = True
165
166    supports_simple_order_by_label = True
167
168    tuple_in_values = False
169
170    connection_characteristics = util.immutabledict(
171        {
172            "isolation_level": characteristics.IsolationLevelCharacteristic(),
173            "logging_token": characteristics.LoggingTokenCharacteristic(),
174        }
175    )
176
177    engine_config_types: Mapping[str, Any] = util.immutabledict(
178        {
179            "pool_timeout": util.asint,
180            "echo": util.bool_or_str("debug"),
181            "echo_pool": util.bool_or_str("debug"),
182            "pool_recycle": util.asint,
183            "pool_size": util.asint,
184            "max_overflow": util.asint,
185            "future": util.asbool,
186        }
187    )
188
189    # if the NUMERIC type
190    # returns decimal.Decimal.
191    # *not* the FLOAT type however.
192    supports_native_decimal = False
193
194    name = "default"
195
196    # length at which to truncate
197    # any identifier.
198    max_identifier_length = 9999
199    _user_defined_max_identifier_length: Optional[int] = None
200
201    isolation_level: Optional[str] = None
202
203    # sub-categories of max_identifier_length.
204    # currently these accommodate for MySQL which allows alias names
205    # of 255 but DDL names only of 64.
206    max_index_name_length: Optional[int] = None
207    max_constraint_name_length: Optional[int] = None
208
209    supports_sane_rowcount = True
210    supports_sane_multi_rowcount = True
211    colspecs: MutableMapping[Type[TypeEngine[Any]], Type[TypeEngine[Any]]] = {}
212    default_paramstyle = "named"
213
214    supports_default_values = False
215    """dialect supports INSERT... DEFAULT VALUES syntax"""
216
217    supports_default_metavalue = False
218    """dialect supports INSERT... VALUES (DEFAULT) syntax"""
219
220    default_metavalue_token = "DEFAULT"
221    """for INSERT... VALUES (DEFAULT) syntax, the token to put in the
222    parenthesis."""
223
224    # not sure if this is a real thing but the compiler will deliver it
225    # if this is the only flag enabled.
226    supports_empty_insert = True
227    """dialect supports INSERT () VALUES ()"""
228
229    supports_multivalues_insert = False
230
231    use_insertmanyvalues: bool = False
232
233    use_insertmanyvalues_wo_returning: bool = False
234
235    insertmanyvalues_implicit_sentinel: InsertmanyvaluesSentinelOpts = (
236        InsertmanyvaluesSentinelOpts.NOT_SUPPORTED
237    )
238
239    insertmanyvalues_page_size: int = 1000
240    insertmanyvalues_max_parameters = 32700
241
242    supports_is_distinct_from = True
243
244    supports_server_side_cursors = False
245
246    server_side_cursors = False
247
248    # extra record-level locking features (#4860)
249    supports_for_update_of = False
250
251    server_version_info = None
252
253    default_schema_name: Optional[str] = None
254
255    # indicates symbol names are
256    # UPPERCASEd if they are case insensitive
257    # within the database.
258    # if this is True, the methods normalize_name()
259    # and denormalize_name() must be provided.
260    requires_name_normalize = False
261
262    is_async = False
263
264    has_terminate = False
265
266    # TODO: this is not to be part of 2.0.  implement rudimentary binary
267    # literals for SQLite, PostgreSQL, MySQL only within
268    # _Binary.literal_processor
269    _legacy_binary_type_literal_encoding = "utf-8"
270
271    @util.deprecated_params(
272        empty_in_strategy=(
273            "1.4",
274            "The :paramref:`_sa.create_engine.empty_in_strategy` keyword is "
275            "deprecated, and no longer has any effect.  All IN expressions "
276            "are now rendered using "
277            'the "expanding parameter" strategy which renders a set of bound'
278            'expressions, or an "empty set" SELECT, at statement execution'
279            "time.",
280        ),
281        server_side_cursors=(
282            "1.4",
283            "The :paramref:`_sa.create_engine.server_side_cursors` parameter "
284            "is deprecated and will be removed in a future release.  Please "
285            "use the "
286            ":paramref:`_engine.Connection.execution_options.stream_results` "
287            "parameter.",
288        ),
289    )
290    def __init__(
291        self,
292        paramstyle: Optional[_ParamStyle] = None,
293        isolation_level: Optional[IsolationLevel] = None,
294        dbapi: Optional[ModuleType] = None,
295        implicit_returning: Literal[True] = True,
296        supports_native_boolean: Optional[bool] = None,
297        max_identifier_length: Optional[int] = None,
298        label_length: Optional[int] = None,
299        insertmanyvalues_page_size: Union[_NoArg, int] = _NoArg.NO_ARG,
300        use_insertmanyvalues: Optional[bool] = None,
301        # util.deprecated_params decorator cannot render the
302        # Linting.NO_LINTING constant
303        compiler_linting: Linting = int(compiler.NO_LINTING),  # type: ignore
304        server_side_cursors: bool = False,
305        **kwargs: Any,
306    ):
307        if server_side_cursors:
308            if not self.supports_server_side_cursors:
309                raise exc.ArgumentError(
310                    "Dialect %s does not support server side cursors" % self
311                )
312            else:
313                self.server_side_cursors = True
314
315        if getattr(self, "use_setinputsizes", False):
316            util.warn_deprecated(
317                "The dialect-level use_setinputsizes attribute is "
318                "deprecated.  Please use "
319                "bind_typing = BindTyping.SETINPUTSIZES",
320                "2.0",
321            )
322            self.bind_typing = interfaces.BindTyping.SETINPUTSIZES
323
324        self.positional = False
325        self._ischema = None
326
327        self.dbapi = dbapi
328
329        if paramstyle is not None:
330            self.paramstyle = paramstyle
331        elif self.dbapi is not None:
332            self.paramstyle = self.dbapi.paramstyle
333        else:
334            self.paramstyle = self.default_paramstyle
335        self.positional = self.paramstyle in (
336            "qmark",
337            "format",
338            "numeric",
339            "numeric_dollar",
340        )
341        self.identifier_preparer = self.preparer(self)
342        self._on_connect_isolation_level = isolation_level
343
344        legacy_tt_callable = getattr(self, "type_compiler", None)
345        if legacy_tt_callable is not None:
346            tt_callable = cast(
347                Type[compiler.GenericTypeCompiler],
348                self.type_compiler,
349            )
350        else:
351            tt_callable = self.type_compiler_cls
352
353        self.type_compiler_instance = self.type_compiler = tt_callable(self)
354
355        if supports_native_boolean is not None:
356            self.supports_native_boolean = supports_native_boolean
357
358        self._user_defined_max_identifier_length = max_identifier_length
359        if self._user_defined_max_identifier_length:
360            self.max_identifier_length = (
361                self._user_defined_max_identifier_length
362            )
363        self.label_length = label_length
364        self.compiler_linting = compiler_linting
365
366        if use_insertmanyvalues is not None:
367            self.use_insertmanyvalues = use_insertmanyvalues
368
369        if insertmanyvalues_page_size is not _NoArg.NO_ARG:
370            self.insertmanyvalues_page_size = insertmanyvalues_page_size
371
372    @property
373    @util.deprecated(
374        "2.0",
375        "full_returning is deprecated, please use insert_returning, "
376        "update_returning, delete_returning",
377    )
378    def full_returning(self):
379        return (
380            self.insert_returning
381            and self.update_returning
382            and self.delete_returning
383        )
384
385    @util.memoized_property
386    def insert_executemany_returning(self):
387        """Default implementation for insert_executemany_returning, if not
388        otherwise overridden by the specific dialect.
389
390        The default dialect determines "insert_executemany_returning" is
391        available if the dialect in use has opted into using the
392        "use_insertmanyvalues" feature. If they haven't opted into that, then
393        this attribute is False, unless the dialect in question overrides this
394        and provides some other implementation (such as the Oracle dialect).
395
396        """
397        return self.insert_returning and self.use_insertmanyvalues
398
399    @util.memoized_property
400    def insert_executemany_returning_sort_by_parameter_order(self):
401        """Default implementation for
402        insert_executemany_returning_deterministic_order, if not otherwise
403        overridden by the specific dialect.
404
405        The default dialect determines "insert_executemany_returning" can have
406        deterministic order only if the dialect in use has opted into using the
407        "use_insertmanyvalues" feature, which implements deterministic ordering
408        using client side sentinel columns only by default.  The
409        "insertmanyvalues" feature also features alternate forms that can
410        use server-generated PK values as "sentinels", but those are only
411        used if the :attr:`.Dialect.insertmanyvalues_implicit_sentinel`
412        bitflag enables those alternate SQL forms, which are disabled
413        by default.
414
415        If the dialect in use hasn't opted into that, then this attribute is
416        False, unless the dialect in question overrides this and provides some
417        other implementation (such as the Oracle dialect).
418
419        """
420        return self.insert_returning and self.use_insertmanyvalues
421
422    update_executemany_returning = False
423    delete_executemany_returning = False
424
425    @util.memoized_property
426    def loaded_dbapi(self) -> ModuleType:
427        if self.dbapi is None:
428            raise exc.InvalidRequestError(
429                f"Dialect {self} does not have a Python DBAPI established "
430                "and cannot be used for actual database interaction"
431            )
432        return self.dbapi
433
434    @util.memoized_property
435    def _bind_typing_render_casts(self):
436        return self.bind_typing is interfaces.BindTyping.RENDER_CASTS
437
438    def _ensure_has_table_connection(self, arg):
439        if not isinstance(arg, Connection):
440            raise exc.ArgumentError(
441                "The argument passed to Dialect.has_table() should be a "
442                "%s, got %s. "
443                "Additionally, the Dialect.has_table() method is for "
444                "internal dialect "
445                "use only; please use "
446                "``inspect(some_engine).has_table(<tablename>>)`` "
447                "for public API use." % (Connection, type(arg))
448            )
449
450    @util.memoized_property
451    def _supports_statement_cache(self):
452        ssc = self.__class__.__dict__.get("supports_statement_cache", None)
453        if ssc is None:
454            util.warn(
455                "Dialect %s:%s will not make use of SQL compilation caching "
456                "as it does not set the 'supports_statement_cache' attribute "
457                "to ``True``.  This can have "
458                "significant performance implications including some "
459                "performance degradations in comparison to prior SQLAlchemy "
460                "versions.  Dialect maintainers should seek to set this "
461                "attribute to True after appropriate development and testing "
462                "for SQLAlchemy 1.4 caching support.   Alternatively, this "
463                "attribute may be set to False which will disable this "
464                "warning." % (self.name, self.driver),
465                code="cprf",
466            )
467
468        return bool(ssc)
469
470    @util.memoized_property
471    def _type_memos(self):
472        return weakref.WeakKeyDictionary()
473
474    @property
475    def dialect_description(self):
476        return self.name + "+" + self.driver
477
478    @property
479    def supports_sane_rowcount_returning(self):
480        """True if this dialect supports sane rowcount even if RETURNING is
481        in use.
482
483        For dialects that don't support RETURNING, this is synonymous with
484        ``supports_sane_rowcount``.
485
486        """
487        return self.supports_sane_rowcount
488
489    @classmethod
490    def get_pool_class(cls, url: URL) -> Type[Pool]:
491        return getattr(cls, "poolclass", pool.QueuePool)
492
493    def get_dialect_pool_class(self, url: URL) -> Type[Pool]:
494        return self.get_pool_class(url)
495
496    @classmethod
497    def load_provisioning(cls):
498        package = ".".join(cls.__module__.split(".")[0:-1])
499        try:
500            __import__(package + ".provision")
501        except ImportError:
502            pass
503
504    def _builtin_onconnect(self) -> Optional[_ListenerFnType]:
505        if self._on_connect_isolation_level is not None:
506
507            def builtin_connect(dbapi_conn, conn_rec):
508                self._assert_and_set_isolation_level(
509                    dbapi_conn, self._on_connect_isolation_level
510                )
511
512            return builtin_connect
513        else:
514            return None
515
516    def initialize(self, connection):
517        try:
518            self.server_version_info = self._get_server_version_info(
519                connection
520            )
521        except NotImplementedError:
522            self.server_version_info = None
523        try:
524            self.default_schema_name = self._get_default_schema_name(
525                connection
526            )
527        except NotImplementedError:
528            self.default_schema_name = None
529
530        try:
531            self.default_isolation_level = self.get_default_isolation_level(
532                connection.connection.dbapi_connection
533            )
534        except NotImplementedError:
535            self.default_isolation_level = None
536
537        if not self._user_defined_max_identifier_length:
538            max_ident_length = self._check_max_identifier_length(connection)
539            if max_ident_length:
540                self.max_identifier_length = max_ident_length
541
542        if (
543            self.label_length
544            and self.label_length > self.max_identifier_length
545        ):
546            raise exc.ArgumentError(
547                "Label length of %d is greater than this dialect's"
548                " maximum identifier length of %d"
549                % (self.label_length, self.max_identifier_length)
550            )
551
552    def on_connect(self):
553        # inherits the docstring from interfaces.Dialect.on_connect
554        return None
555
556    def _check_max_identifier_length(self, connection):
557        """Perform a connection / server version specific check to determine
558        the max_identifier_length.
559
560        If the dialect's class level max_identifier_length should be used,
561        can return None.
562
563        .. versionadded:: 1.3.9
564
565        """
566        return None
567
568    def get_default_isolation_level(self, dbapi_conn):
569        """Given a DBAPI connection, return its isolation level, or
570        a default isolation level if one cannot be retrieved.
571
572        May be overridden by subclasses in order to provide a
573        "fallback" isolation level for databases that cannot reliably
574        retrieve the actual isolation level.
575
576        By default, calls the :meth:`_engine.Interfaces.get_isolation_level`
577        method, propagating any exceptions raised.
578
579        .. versionadded:: 1.3.22
580
581        """
582        return self.get_isolation_level(dbapi_conn)
583
584    def type_descriptor(self, typeobj):
585        """Provide a database-specific :class:`.TypeEngine` object, given
586        the generic object which comes from the types module.
587
588        This method looks for a dictionary called
589        ``colspecs`` as a class or instance-level variable,
590        and passes on to :func:`_types.adapt_type`.
591
592        """
593        return type_api.adapt_type(typeobj, self.colspecs)
594
595    def has_index(self, connection, table_name, index_name, schema=None, **kw):
596        if not self.has_table(connection, table_name, schema=schema, **kw):
597            return False
598        for idx in self.get_indexes(
599            connection, table_name, schema=schema, **kw
600        ):
601            if idx["name"] == index_name:
602                return True
603        else:
604            return False
605
606    def has_schema(
607        self, connection: Connection, schema_name: str, **kw: Any
608    ) -> bool:
609        return schema_name in self.get_schema_names(connection, **kw)
610
611    def validate_identifier(self, ident):
612        if len(ident) > self.max_identifier_length:
613            raise exc.IdentifierError(
614                "Identifier '%s' exceeds maximum length of %d characters"
615                % (ident, self.max_identifier_length)
616            )
617
618    def connect(self, *cargs, **cparams):
619        # inherits the docstring from interfaces.Dialect.connect
620        return self.loaded_dbapi.connect(*cargs, **cparams)
621
622    def create_connect_args(self, url):
623        # inherits the docstring from interfaces.Dialect.create_connect_args
624        opts = url.translate_connect_args()
625        opts.update(url.query)
626        return ([], opts)
627
628    def set_engine_execution_options(
629        self, engine: Engine, opts: Mapping[str, Any]
630    ) -> None:
631        supported_names = set(self.connection_characteristics).intersection(
632            opts
633        )
634        if supported_names:
635            characteristics: Mapping[str, Any] = util.immutabledict(
636                (name, opts[name]) for name in supported_names
637            )
638
639            @event.listens_for(engine, "engine_connect")
640            def set_connection_characteristics(connection):
641                self._set_connection_characteristics(
642                    connection, characteristics
643                )
644
645    def set_connection_execution_options(
646        self, connection: Connection, opts: Mapping[str, Any]
647    ) -> None:
648        supported_names = set(self.connection_characteristics).intersection(
649            opts
650        )
651        if supported_names:
652            characteristics: Mapping[str, Any] = util.immutabledict(
653                (name, opts[name]) for name in supported_names
654            )
655            self._set_connection_characteristics(connection, characteristics)
656
657    def _set_connection_characteristics(self, connection, characteristics):
658        characteristic_values = [
659            (name, self.connection_characteristics[name], value)
660            for name, value in characteristics.items()
661        ]
662
663        if connection.in_transaction():
664            trans_objs = [
665                (name, obj)
666                for name, obj, _ in characteristic_values
667                if obj.transactional
668            ]
669            if trans_objs:
670                raise exc.InvalidRequestError(
671                    "This connection has already initialized a SQLAlchemy "
672                    "Transaction() object via begin() or autobegin; "
673                    "%s may not be altered unless rollback() or commit() "
674                    "is called first."
675                    % (", ".join(name for name, obj in trans_objs))
676                )
677
678        dbapi_connection = connection.connection.dbapi_connection
679        for _, characteristic, value in characteristic_values:
680            characteristic.set_connection_characteristic(
681                self, connection, dbapi_connection, value
682            )
683        connection.connection._connection_record.finalize_callback.append(
684            functools.partial(self._reset_characteristics, characteristics)
685        )
686
687    def _reset_characteristics(self, characteristics, dbapi_connection):
688        for characteristic_name in characteristics:
689            characteristic = self.connection_characteristics[
690                characteristic_name
691            ]
692            characteristic.reset_characteristic(self, dbapi_connection)
693
694    def do_begin(self, dbapi_connection):
695        pass
696
697    def do_rollback(self, dbapi_connection):
698        dbapi_connection.rollback()
699
700    def do_commit(self, dbapi_connection):
701        dbapi_connection.commit()
702
703    def do_terminate(self, dbapi_connection):
704        self.do_close(dbapi_connection)
705
706    def do_close(self, dbapi_connection):
707        dbapi_connection.close()
708
709    @util.memoized_property
710    def _dialect_specific_select_one(self):
711        return str(expression.select(1).compile(dialect=self))
712
713    def _do_ping_w_event(self, dbapi_connection: DBAPIConnection) -> bool:
714        try:
715            return self.do_ping(dbapi_connection)
716        except self.loaded_dbapi.Error as err:
717            is_disconnect = self.is_disconnect(err, dbapi_connection, None)
718
719            if self._has_events:
720                try:
721                    Connection._handle_dbapi_exception_noconnection(
722                        err,
723                        self,
724                        is_disconnect=is_disconnect,
725                        invalidate_pool_on_disconnect=False,
726                        is_pre_ping=True,
727                    )
728                except exc.StatementError as new_err:
729                    is_disconnect = new_err.connection_invalidated
730
731            if is_disconnect:
732                return False
733            else:
734                raise
735
736    def do_ping(self, dbapi_connection: DBAPIConnection) -> bool:
737        cursor = None
738
739        cursor = dbapi_connection.cursor()
740        try:
741            cursor.execute(self._dialect_specific_select_one)
742        finally:
743            cursor.close()
744        return True
745
746    def create_xid(self):
747        """Create a random two-phase transaction ID.
748
749        This id will be passed to do_begin_twophase(), do_rollback_twophase(),
750        do_commit_twophase().  Its format is unspecified.
751        """
752
753        return "_sa_%032x" % random.randint(0, 2**128)
754
755    def do_savepoint(self, connection, name):
756        connection.execute(expression.SavepointClause(name))
757
758    def do_rollback_to_savepoint(self, connection, name):
759        connection.execute(expression.RollbackToSavepointClause(name))
760
761    def do_release_savepoint(self, connection, name):
762        connection.execute(expression.ReleaseSavepointClause(name))
763
764    def _deliver_insertmanyvalues_batches(
765        self, cursor, statement, parameters, generic_setinputsizes, context
766    ):
767        context = cast(DefaultExecutionContext, context)
768        compiled = cast(SQLCompiler, context.compiled)
769
770        _composite_sentinel_proc: Sequence[
771            Optional[_ResultProcessorType[Any]]
772        ] = ()
773        _scalar_sentinel_proc: Optional[_ResultProcessorType[Any]] = None
774        _sentinel_proc_initialized: bool = False
775
776        compiled_parameters = context.compiled_parameters
777
778        imv = compiled._insertmanyvalues
779        assert imv is not None
780
781        is_returning: Final[bool] = bool(compiled.effective_returning)
782        batch_size = context.execution_options.get(
783            "insertmanyvalues_page_size", self.insertmanyvalues_page_size
784        )
785
786        if compiled.schema_translate_map:
787            schema_translate_map = context.execution_options.get(
788                "schema_translate_map", {}
789            )
790        else:
791            schema_translate_map = None
792
793        if is_returning:
794            result: Optional[List[Any]] = []
795            context._insertmanyvalues_rows = result
796
797            sort_by_parameter_order = imv.sort_by_parameter_order
798
799        else:
800            sort_by_parameter_order = False
801            result = None
802
803        for imv_batch in compiled._deliver_insertmanyvalues_batches(
804            statement,
805            parameters,
806            compiled_parameters,
807            generic_setinputsizes,
808            batch_size,
809            sort_by_parameter_order,
810            schema_translate_map,
811        ):
812            yield imv_batch
813
814            if is_returning:
815
816                rows = context.fetchall_for_returning(cursor)
817
818                # I would have thought "is_returning: Final[bool]"
819                # would have assured this but pylance thinks not
820                assert result is not None
821
822                if imv.num_sentinel_columns and not imv_batch.is_downgraded:
823                    composite_sentinel = imv.num_sentinel_columns > 1
824                    if imv.implicit_sentinel:
825                        # for implicit sentinel, which is currently single-col
826                        # integer autoincrement, do a simple sort.
827                        assert not composite_sentinel
828                        result.extend(
829                            sorted(rows, key=operator.itemgetter(-1))
830                        )
831                        continue
832
833                    # otherwise, create dictionaries to match up batches
834                    # with parameters
835                    assert imv.sentinel_param_keys
836                    assert imv.sentinel_columns
837
838                    _nsc = imv.num_sentinel_columns
839
840                    if not _sentinel_proc_initialized:
841                        if composite_sentinel:
842                            _composite_sentinel_proc = [
843                                col.type._cached_result_processor(
844                                    self, cursor_desc[1]
845                                )
846                                for col, cursor_desc in zip(
847                                    imv.sentinel_columns,
848                                    cursor.description[-_nsc:],
849                                )
850                            ]
851                        else:
852                            _scalar_sentinel_proc = (
853                                imv.sentinel_columns[0]
854                            ).type._cached_result_processor(
855                                self, cursor.description[-1][1]
856                            )
857                        _sentinel_proc_initialized = True
858
859                    rows_by_sentinel: Union[
860                        Dict[Tuple[Any, ...], Any],
861                        Dict[Any, Any],
862                    ]
863                    if composite_sentinel:
864                        rows_by_sentinel = {
865                            tuple(
866                                (proc(val) if proc else val)
867                                for val, proc in zip(
868                                    row[-_nsc:], _composite_sentinel_proc
869                                )
870                            ): row
871                            for row in rows
872                        }
873                    elif _scalar_sentinel_proc:
874                        rows_by_sentinel = {
875                            _scalar_sentinel_proc(row[-1]): row for row in rows
876                        }
877                    else:
878                        rows_by_sentinel = {row[-1]: row for row in rows}
879
880                    if len(rows_by_sentinel) != len(imv_batch.batch):
881                        # see test_insert_exec.py::
882                        # IMVSentinelTest::test_sentinel_incorrect_rowcount
883                        # for coverage / demonstration
884                        raise exc.InvalidRequestError(
885                            f"Sentinel-keyed result set did not produce "
886                            f"correct number of rows {len(imv_batch.batch)}; "
887                            "produced "
888                            f"{len(rows_by_sentinel)}.  Please ensure the "
889                            "sentinel column is fully unique and populated in "
890                            "all cases."
891                        )
892
893                    try:
894                        ordered_rows = [
895                            rows_by_sentinel[sentinel_keys]
896                            for sentinel_keys in imv_batch.sentinel_values
897                        ]
898                    except KeyError as ke:
899                        # see test_insert_exec.py::
900                        # IMVSentinelTest::test_sentinel_cant_match_keys
901                        # for coverage / demonstration
902                        raise exc.InvalidRequestError(
903                            f"Can't match sentinel values in result set to "
904                            f"parameter sets; key {ke.args[0]!r} was not "
905                            "found. "
906                            "There may be a mismatch between the datatype "
907                            "passed to the DBAPI driver vs. that which it "
908                            "returns in a result row.  Ensure the given "
909                            "Python value matches the expected result type "
910                            "*exactly*, taking care to not rely upon implicit "
911                            "conversions which may occur such as when using "
912                            "strings in place of UUID or integer values, etc. "
913                        ) from ke
914
915                    result.extend(ordered_rows)
916
917                else:
918                    result.extend(rows)
919
920    def do_executemany(self, cursor, statement, parameters, context=None):
921        cursor.executemany(statement, parameters)
922
923    def do_execute(self, cursor, statement, parameters, context=None):
924        cursor.execute(statement, parameters)
925
926    def do_execute_no_params(self, cursor, statement, context=None):
927        cursor.execute(statement)
928
929    def is_disconnect(self, e, connection, cursor):
930        return False
931
932    @util.memoized_instancemethod
933    def _gen_allowed_isolation_levels(self, dbapi_conn):
934        try:
935            raw_levels = list(self.get_isolation_level_values(dbapi_conn))
936        except NotImplementedError:
937            return None
938        else:
939            normalized_levels = [
940                level.replace("_", " ").upper() for level in raw_levels
941            ]
942            if raw_levels != normalized_levels:
943                raise ValueError(
944                    f"Dialect {self.name!r} get_isolation_level_values() "
945                    f"method should return names as UPPERCASE using spaces, "
946                    f"not underscores; got "
947                    f"{sorted(set(raw_levels).difference(normalized_levels))}"
948                )
949            return tuple(normalized_levels)
950
951    def _assert_and_set_isolation_level(self, dbapi_conn, level):
952        level = level.replace("_", " ").upper()
953
954        _allowed_isolation_levels = self._gen_allowed_isolation_levels(
955            dbapi_conn
956        )
957        if (
958            _allowed_isolation_levels
959            and level not in _allowed_isolation_levels
960        ):
961            raise exc.ArgumentError(
962                f"Invalid value {level!r} for isolation_level. "
963                f"Valid isolation levels for {self.name!r} are "
964                f"{', '.join(_allowed_isolation_levels)}"
965            )
966
967        self.set_isolation_level(dbapi_conn, level)
968
969    def reset_isolation_level(self, dbapi_conn):
970        if self._on_connect_isolation_level is not None:
971            assert (
972                self._on_connect_isolation_level == "AUTOCOMMIT"
973                or self._on_connect_isolation_level
974                == self.default_isolation_level
975            )
976            self._assert_and_set_isolation_level(
977                dbapi_conn, self._on_connect_isolation_level
978            )
979        else:
980            assert self.default_isolation_level is not None
981            self._assert_and_set_isolation_level(
982                dbapi_conn,
983                self.default_isolation_level,
984            )
985
986    def normalize_name(self, name):
987        if name is None:
988            return None
989
990        name_lower = name.lower()
991        name_upper = name.upper()
992
993        if name_upper == name_lower:
994            # name has no upper/lower conversion, e.g. non-european characters.
995            # return unchanged
996            return name
997        elif name_upper == name and not (
998            self.identifier_preparer._requires_quotes
999        )(name_lower):
1000            # name is all uppercase and doesn't require quoting; normalize
1001            # to all lower case
1002            return name_lower
1003        elif name_lower == name:
1004            # name is all lower case, which if denormalized means we need to
1005            # force quoting on it
1006            return quoted_name(name, quote=True)
1007        else:
1008            # name is mixed case, means it will be quoted in SQL when used
1009            # later, no normalizes
1010            return name
1011
1012    def denormalize_name(self, name):
1013        if name is None:
1014            return None
1015
1016        name_lower = name.lower()
1017        name_upper = name.upper()
1018
1019        if name_upper == name_lower:
1020            # name has no upper/lower conversion, e.g. non-european characters.
1021            # return unchanged
1022            return name
1023        elif name_lower == name and not (
1024            self.identifier_preparer._requires_quotes
1025        )(name_lower):
1026            name = name_upper
1027        return name
1028
1029    def get_driver_connection(self, connection):
1030        return connection
1031
1032    def _overrides_default(self, method):
1033        return (
1034            getattr(type(self), method).__code__
1035            is not getattr(DefaultDialect, method).__code__
1036        )
1037
1038    def _default_multi_reflect(
1039        self,
1040        single_tbl_method,
1041        connection,
1042        kind,
1043        schema,
1044        filter_names,
1045        scope,
1046        **kw,
1047    ):
1048        names_fns = []
1049        temp_names_fns = []
1050        if ObjectKind.TABLE in kind:
1051            names_fns.append(self.get_table_names)
1052            temp_names_fns.append(self.get_temp_table_names)
1053        if ObjectKind.VIEW in kind:
1054            names_fns.append(self.get_view_names)
1055            temp_names_fns.append(self.get_temp_view_names)
1056        if ObjectKind.MATERIALIZED_VIEW in kind:
1057            names_fns.append(self.get_materialized_view_names)
1058            # no temp materialized view at the moment
1059            # temp_names_fns.append(self.get_temp_materialized_view_names)
1060
1061        unreflectable = kw.pop("unreflectable", {})
1062
1063        if (
1064            filter_names
1065            and scope is ObjectScope.ANY
1066            and kind is ObjectKind.ANY
1067        ):
1068            # if names are given and no qualification on type of table
1069            # (i.e. the Table(..., autoload) case), take the names as given,
1070            # don't run names queries. If a table does not exit
1071            # NoSuchTableError is raised and it's skipped
1072
1073            # this also suits the case for mssql where we can reflect
1074            # individual temp tables but there's no temp_names_fn
1075            names = filter_names
1076        else:
1077            names = []
1078            name_kw = {"schema": schema, **kw}
1079            fns = []
1080            if ObjectScope.DEFAULT in scope:
1081                fns.extend(names_fns)
1082            if ObjectScope.TEMPORARY in scope:
1083                fns.extend(temp_names_fns)
1084
1085            for fn in fns:
1086                try:
1087                    names.extend(fn(connection, **name_kw))
1088                except NotImplementedError:
1089                    pass
1090
1091        if filter_names:
1092            filter_names = set(filter_names)
1093
1094        # iterate over all the tables/views and call the single table method
1095        for table in names:
1096            if not filter_names or table in filter_names:
1097                key = (schema, table)
1098                try:
1099                    yield (
1100                        key,
1101                        single_tbl_method(
1102                            connection, table, schema=schema, **kw
1103                        ),
1104                    )
1105                except exc.UnreflectableTableError as err:
1106                    if key not in unreflectable:
1107                        unreflectable[key] = err
1108                except exc.NoSuchTableError:
1109                    pass
1110
1111    def get_multi_table_options(self, connection, **kw):
1112        return self._default_multi_reflect(
1113            self.get_table_options, connection, **kw
1114        )
1115
1116    def get_multi_columns(self, connection, **kw):
1117        return self._default_multi_reflect(self.get_columns, connection, **kw)
1118
1119    def get_multi_pk_constraint(self, connection, **kw):
1120        return self._default_multi_reflect(
1121            self.get_pk_constraint, connection, **kw
1122        )
1123
1124    def get_multi_foreign_keys(self, connection, **kw):
1125        return self._default_multi_reflect(
1126            self.get_foreign_keys, connection, **kw
1127        )
1128
1129    def get_multi_indexes(self, connection, **kw):
1130        return self._default_multi_reflect(self.get_indexes, connection, **kw)
1131
1132    def get_multi_unique_constraints(self, connection, **kw):
1133        return self._default_multi_reflect(
1134            self.get_unique_constraints, connection, **kw
1135        )
1136
1137    def get_multi_check_constraints(self, connection, **kw):
1138        return self._default_multi_reflect(
1139            self.get_check_constraints, connection, **kw
1140        )
1141
1142    def get_multi_table_comment(self, connection, **kw):
1143        return self._default_multi_reflect(
1144            self.get_table_comment, connection, **kw
1145        )
1146
1147
1148class StrCompileDialect(DefaultDialect):
1149    statement_compiler = compiler.StrSQLCompiler
1150    ddl_compiler = compiler.DDLCompiler
1151    type_compiler_cls = compiler.StrSQLTypeCompiler
1152    preparer = compiler.IdentifierPreparer
1153
1154    insert_returning = True
1155    update_returning = True
1156    delete_returning = True
1157
1158    supports_statement_cache = True
1159
1160    supports_identity_columns = True
1161
1162    supports_sequences = True
1163    sequences_optional = True
1164    preexecute_autoincrement_sequences = False
1165
1166    supports_native_boolean = True
1167
1168    supports_multivalues_insert = True
1169    supports_simple_order_by_label = True
1170
1171
1172class DefaultExecutionContext(ExecutionContext):
1173    isinsert = False
1174    isupdate = False
1175    isdelete = False
1176    is_crud = False
1177    is_text = False
1178    isddl = False
1179
1180    execute_style: ExecuteStyle = ExecuteStyle.EXECUTE
1181
1182    compiled: Optional[Compiled] = None
1183    result_column_struct: Optional[
1184        Tuple[List[ResultColumnsEntry], bool, bool, bool, bool]
1185    ] = None
1186    returned_default_rows: Optional[Sequence[Row[Any]]] = None
1187
1188    execution_options: _ExecuteOptions = util.EMPTY_DICT
1189
1190    cursor_fetch_strategy = _cursor._DEFAULT_FETCH
1191
1192    invoked_statement: Optional[Executable] = None
1193
1194    _is_implicit_returning = False
1195    _is_explicit_returning = False
1196    _is_supplemental_returning = False
1197    _is_server_side = False
1198
1199    _soft_closed = False
1200

Showing the first 1,200 of 2349 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai