Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
bulk_persistence.py2056 linesDownload Raw Back to orm
1# orm/bulk_persistence.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
10"""additional ORM persistence classes related to "bulk" operations,
11specifically outside of the flush() process.
12
13"""
14
15from __future__ import annotations
16
17from typing import Any
18from typing import cast
19from typing import Dict
20from typing import Iterable
21from typing import Optional
22from typing import overload
23from typing import TYPE_CHECKING
24from typing import TypeVar
25from typing import Union
26
27from . import attributes
28from . import context
29from . import evaluator
30from . import exc as orm_exc
31from . import loading
32from . import persistence
33from .base import NO_VALUE
34from .context import AbstractORMCompileState
35from .context import FromStatement
36from .context import ORMFromStatementCompileState
37from .context import QueryContext
38from .. import exc as sa_exc
39from .. import util
40from ..engine import Dialect
41from ..engine import result as _result
42from ..sql import coercions
43from ..sql import dml
44from ..sql import expression
45from ..sql import roles
46from ..sql import select
47from ..sql import sqltypes
48from ..sql.base import _entity_namespace_key
49from ..sql.base import CompileState
50from ..sql.base import Options
51from ..sql.dml import DeleteDMLState
52from ..sql.dml import InsertDMLState
53from ..sql.dml import UpdateDMLState
54from ..util import EMPTY_DICT
55from ..util.typing import Literal
56
57if TYPE_CHECKING:
58    from ._typing import DMLStrategyArgument
59    from ._typing import OrmExecuteOptionsParameter
60    from ._typing import SynchronizeSessionArgument
61    from .mapper import Mapper
62    from .session import _BindArguments
63    from .session import ORMExecuteState
64    from .session import Session
65    from .session import SessionTransaction
66    from .state import InstanceState
67    from ..engine import Connection
68    from ..engine import cursor
69    from ..engine.interfaces import _CoreAnyExecuteParams
70
71_O = TypeVar("_O", bound=object)
72
73
74@overload
75def _bulk_insert(
76    mapper: Mapper[_O],
77    mappings: Union[Iterable[InstanceState[_O]], Iterable[Dict[str, Any]]],
78    session_transaction: SessionTransaction,
79    *,
80    isstates: bool,
81    return_defaults: bool,
82    render_nulls: bool,
83    use_orm_insert_stmt: Literal[None] = ...,
84    execution_options: Optional[OrmExecuteOptionsParameter] = ...,
85) -> None: ...
86
87
88@overload
89def _bulk_insert(
90    mapper: Mapper[_O],
91    mappings: Union[Iterable[InstanceState[_O]], Iterable[Dict[str, Any]]],
92    session_transaction: SessionTransaction,
93    *,
94    isstates: bool,
95    return_defaults: bool,
96    render_nulls: bool,
97    use_orm_insert_stmt: Optional[dml.Insert] = ...,
98    execution_options: Optional[OrmExecuteOptionsParameter] = ...,
99) -> cursor.CursorResult[Any]: ...
100
101
102def _bulk_insert(
103    mapper: Mapper[_O],
104    mappings: Union[Iterable[InstanceState[_O]], Iterable[Dict[str, Any]]],
105    session_transaction: SessionTransaction,
106    *,
107    isstates: bool,
108    return_defaults: bool,
109    render_nulls: bool,
110    use_orm_insert_stmt: Optional[dml.Insert] = None,
111    execution_options: Optional[OrmExecuteOptionsParameter] = None,
112) -> Optional[cursor.CursorResult[Any]]:
113    base_mapper = mapper.base_mapper
114
115    if session_transaction.session.connection_callable:
116        raise NotImplementedError(
117            "connection_callable / per-instance sharding "
118            "not supported in bulk_insert()"
119        )
120
121    if isstates:
122        if return_defaults:
123            states = [(state, state.dict) for state in mappings]
124            mappings = [dict_ for (state, dict_) in states]
125        else:
126            mappings = [state.dict for state in mappings]
127    else:
128        mappings = [dict(m) for m in mappings]
129        _expand_composites(mapper, mappings)
130
131    connection = session_transaction.connection(base_mapper)
132
133    return_result: Optional[cursor.CursorResult[Any]] = None
134
135    mappers_to_run = [
136        (table, mp)
137        for table, mp in base_mapper._sorted_tables.items()
138        if table in mapper._pks_by_table
139    ]
140
141    if return_defaults:
142        # not used by new-style bulk inserts, only used for legacy
143        bookkeeping = True
144    elif len(mappers_to_run) > 1:
145        # if we have more than one table, mapper to run where we will be
146        # either horizontally splicing, or copying values between tables,
147        # we need the "bookkeeping" / deterministic returning order
148        bookkeeping = True
149    else:
150        bookkeeping = False
151
152    for table, super_mapper in mappers_to_run:
153        # find bindparams in the statement. For bulk, we don't really know if
154        # a key in the params applies to a different table since we are
155        # potentially inserting for multiple tables here; looking at the
156        # bindparam() is a lot more direct.   in most cases this will
157        # use _generate_cache_key() which is memoized, although in practice
158        # the ultimate statement that's executed is probably not the same
159        # object so that memoization might not matter much.
160        extra_bp_names = (
161            [
162                b.key
163                for b in use_orm_insert_stmt._get_embedded_bindparams()
164                if b.key in mappings[0]
165            ]
166            if use_orm_insert_stmt is not None
167            else ()
168        )
169
170        records = (
171            (
172                None,
173                state_dict,
174                params,
175                mapper,
176                connection,
177                value_params,
178                has_all_pks,
179                has_all_defaults,
180            )
181            for (
182                state,
183                state_dict,
184                params,
185                mp,
186                conn,
187                value_params,
188                has_all_pks,
189                has_all_defaults,
190            ) in persistence._collect_insert_commands(
191                table,
192                ((None, mapping, mapper, connection) for mapping in mappings),
193                bulk=True,
194                return_defaults=bookkeeping,
195                render_nulls=render_nulls,
196                include_bulk_keys=extra_bp_names,
197            )
198        )
199
200        result = persistence._emit_insert_statements(
201            base_mapper,
202            None,
203            super_mapper,
204            table,
205            records,
206            bookkeeping=bookkeeping,
207            use_orm_insert_stmt=use_orm_insert_stmt,
208            execution_options=execution_options,
209        )
210        if use_orm_insert_stmt is not None:
211            if not use_orm_insert_stmt._returning or return_result is None:
212                return_result = result
213            elif result.returns_rows:
214                assert bookkeeping
215                return_result = return_result.splice_horizontally(result)
216
217    if return_defaults and isstates:
218        identity_cls = mapper._identity_class
219        identity_props = [p.key for p in mapper._identity_key_props]
220        for state, dict_ in states:
221            state.key = (
222                identity_cls,
223                tuple([dict_[key] for key in identity_props]),
224                None,
225            )
226
227    if use_orm_insert_stmt is not None:
228        assert return_result is not None
229        return return_result
230
231
232@overload
233def _bulk_update(
234    mapper: Mapper[Any],
235    mappings: Union[Iterable[InstanceState[_O]], Iterable[Dict[str, Any]]],
236    session_transaction: SessionTransaction,
237    *,
238    isstates: bool,
239    update_changed_only: bool,
240    use_orm_update_stmt: Literal[None] = ...,
241    enable_check_rowcount: bool = True,
242) -> None: ...
243
244
245@overload
246def _bulk_update(
247    mapper: Mapper[Any],
248    mappings: Union[Iterable[InstanceState[_O]], Iterable[Dict[str, Any]]],
249    session_transaction: SessionTransaction,
250    *,
251    isstates: bool,
252    update_changed_only: bool,
253    use_orm_update_stmt: Optional[dml.Update] = ...,
254    enable_check_rowcount: bool = True,
255) -> _result.Result[Any]: ...
256
257
258def _bulk_update(
259    mapper: Mapper[Any],
260    mappings: Union[Iterable[InstanceState[_O]], Iterable[Dict[str, Any]]],
261    session_transaction: SessionTransaction,
262    *,
263    isstates: bool,
264    update_changed_only: bool,
265    use_orm_update_stmt: Optional[dml.Update] = None,
266    enable_check_rowcount: bool = True,
267) -> Optional[_result.Result[Any]]:
268    base_mapper = mapper.base_mapper
269
270    search_keys = mapper._primary_key_propkeys
271    if mapper._version_id_prop:
272        search_keys = {mapper._version_id_prop.key}.union(search_keys)
273
274    def _changed_dict(mapper, state):
275        return {
276            k: v
277            for k, v in state.dict.items()
278            if k in state.committed_state or k in search_keys
279        }
280
281    if isstates:
282        if update_changed_only:
283            mappings = [_changed_dict(mapper, state) for state in mappings]
284        else:
285            mappings = [state.dict for state in mappings]
286    else:
287        mappings = [dict(m) for m in mappings]
288        _expand_composites(mapper, mappings)
289
290    if session_transaction.session.connection_callable:
291        raise NotImplementedError(
292            "connection_callable / per-instance sharding "
293            "not supported in bulk_update()"
294        )
295
296    connection = session_transaction.connection(base_mapper)
297
298    # find bindparams in the statement. see _bulk_insert for similar
299    # notes for the insert case
300    extra_bp_names = (
301        [
302            b.key
303            for b in use_orm_update_stmt._get_embedded_bindparams()
304            if b.key in mappings[0]
305        ]
306        if use_orm_update_stmt is not None
307        else ()
308    )
309
310    for table, super_mapper in base_mapper._sorted_tables.items():
311        if not mapper.isa(super_mapper) or table not in mapper._pks_by_table:
312            continue
313
314        records = persistence._collect_update_commands(
315            None,
316            table,
317            (
318                (
319                    None,
320                    mapping,
321                    mapper,
322                    connection,
323                    (
324                        mapping[mapper._version_id_prop.key]
325                        if mapper._version_id_prop
326                        else None
327                    ),
328                )
329                for mapping in mappings
330            ),
331            bulk=True,
332            use_orm_update_stmt=use_orm_update_stmt,
333            include_bulk_keys=extra_bp_names,
334        )
335        persistence._emit_update_statements(
336            base_mapper,
337            None,
338            super_mapper,
339            table,
340            records,
341            bookkeeping=False,
342            use_orm_update_stmt=use_orm_update_stmt,
343            enable_check_rowcount=enable_check_rowcount,
344        )
345
346    if use_orm_update_stmt is not None:
347        return _result.null_result()
348
349
350def _expand_composites(mapper, mappings):
351    composite_attrs = mapper.composites
352    if not composite_attrs:
353        return
354
355    composite_keys = set(composite_attrs.keys())
356    populators = {
357        key: composite_attrs[key]._populate_composite_bulk_save_mappings_fn()
358        for key in composite_keys
359    }
360    for mapping in mappings:
361        for key in composite_keys.intersection(mapping):
362            populators[key](mapping)
363
364
365class ORMDMLState(AbstractORMCompileState):
366    is_dml_returning = True
367    from_statement_ctx: Optional[ORMFromStatementCompileState] = None
368
369    @classmethod
370    def _get_orm_crud_kv_pairs(
371        cls, mapper, statement, kv_iterator, needs_to_be_cacheable
372    ):
373        core_get_crud_kv_pairs = UpdateDMLState._get_crud_kv_pairs
374
375        for k, v in kv_iterator:
376            k = coercions.expect(roles.DMLColumnRole, k)
377
378            if isinstance(k, str):
379                desc = _entity_namespace_key(mapper, k, default=NO_VALUE)
380                if desc is NO_VALUE:
381                    yield (
382                        coercions.expect(roles.DMLColumnRole, k),
383                        (
384                            coercions.expect(
385                                roles.ExpressionElementRole,
386                                v,
387                                type_=sqltypes.NullType(),
388                                is_crud=True,
389                            )
390                            if needs_to_be_cacheable
391                            else v
392                        ),
393                    )
394                else:
395                    yield from core_get_crud_kv_pairs(
396                        statement,
397                        desc._bulk_update_tuples(v),
398                        needs_to_be_cacheable,
399                    )
400            elif "entity_namespace" in k._annotations:
401                k_anno = k._annotations
402                attr = _entity_namespace_key(
403                    k_anno["entity_namespace"], k_anno["proxy_key"]
404                )
405                yield from core_get_crud_kv_pairs(
406                    statement,
407                    attr._bulk_update_tuples(v),
408                    needs_to_be_cacheable,
409                )
410            else:
411                yield (
412                    k,
413                    (
414                        v
415                        if not needs_to_be_cacheable
416                        else coercions.expect(
417                            roles.ExpressionElementRole,
418                            v,
419                            type_=sqltypes.NullType(),
420                            is_crud=True,
421                        )
422                    ),
423                )
424
425    @classmethod
426    def _get_multi_crud_kv_pairs(cls, statement, kv_iterator):
427        plugin_subject = statement._propagate_attrs["plugin_subject"]
428
429        if not plugin_subject or not plugin_subject.mapper:
430            return UpdateDMLState._get_multi_crud_kv_pairs(
431                statement, kv_iterator
432            )
433
434        return [
435            dict(
436                cls._get_orm_crud_kv_pairs(
437                    plugin_subject.mapper, statement, value_dict.items(), False
438                )
439            )
440            for value_dict in kv_iterator
441        ]
442
443    @classmethod
444    def _get_crud_kv_pairs(cls, statement, kv_iterator, needs_to_be_cacheable):
445        assert (
446            needs_to_be_cacheable
447        ), "no test coverage for needs_to_be_cacheable=False"
448
449        plugin_subject = statement._propagate_attrs["plugin_subject"]
450
451        if not plugin_subject or not plugin_subject.mapper:
452            return UpdateDMLState._get_crud_kv_pairs(
453                statement, kv_iterator, needs_to_be_cacheable
454            )
455
456        return list(
457            cls._get_orm_crud_kv_pairs(
458                plugin_subject.mapper,
459                statement,
460                kv_iterator,
461                needs_to_be_cacheable,
462            )
463        )
464
465    @classmethod
466    def get_entity_description(cls, statement):
467        ext_info = statement.table._annotations["parententity"]
468        mapper = ext_info.mapper
469        if ext_info.is_aliased_class:
470            _label_name = ext_info.name
471        else:
472            _label_name = mapper.class_.__name__
473
474        return {
475            "name": _label_name,
476            "type": mapper.class_,
477            "expr": ext_info.entity,
478            "entity": ext_info.entity,
479            "table": mapper.local_table,
480        }
481
482    @classmethod
483    def get_returning_column_descriptions(cls, statement):
484        def _ent_for_col(c):
485            return c._annotations.get("parententity", None)
486
487        def _attr_for_col(c, ent):
488            if ent is None:
489                return c
490            proxy_key = c._annotations.get("proxy_key", None)
491            if not proxy_key:
492                return c
493            else:
494                return getattr(ent.entity, proxy_key, c)
495
496        return [
497            {
498                "name": c.key,
499                "type": c.type,
500                "expr": _attr_for_col(c, ent),
501                "aliased": ent.is_aliased_class,
502                "entity": ent.entity,
503            }
504            for c, ent in [
505                (c, _ent_for_col(c)) for c in statement._all_selected_columns
506            ]
507        ]
508
509    def _setup_orm_returning(
510        self,
511        compiler,
512        orm_level_statement,
513        dml_level_statement,
514        dml_mapper,
515        *,
516        use_supplemental_cols=True,
517    ):
518        """establish ORM column handlers for an INSERT, UPDATE, or DELETE
519        which uses explicit returning().
520
521        called within compilation level create_for_statement.
522
523        The _return_orm_returning() method then receives the Result
524        after the statement was executed, and applies ORM loading to the
525        state that we first established here.
526
527        """
528
529        if orm_level_statement._returning:
530            fs = FromStatement(
531                orm_level_statement._returning,
532                dml_level_statement,
533                _adapt_on_names=False,
534            )
535            fs = fs.execution_options(**orm_level_statement._execution_options)
536            fs = fs.options(*orm_level_statement._with_options)
537            self.select_statement = fs
538            self.from_statement_ctx = fsc = (
539                ORMFromStatementCompileState.create_for_statement(fs, compiler)
540            )
541            fsc.setup_dml_returning_compile_state(dml_mapper)
542
543            dml_level_statement = dml_level_statement._generate()
544            dml_level_statement._returning = ()
545
546            cols_to_return = [c for c in fsc.primary_columns if c is not None]
547
548            # since we are splicing result sets together, make sure there
549            # are columns of some kind returned in each result set
550            if not cols_to_return:
551                cols_to_return.extend(dml_mapper.primary_key)
552
553            if use_supplemental_cols:
554                dml_level_statement = dml_level_statement.return_defaults(
555                    # this is a little weird looking, but by passing
556                    # primary key as the main list of cols, this tells
557                    # return_defaults to omit server-default cols (and
558                    # actually all cols, due to some weird thing we should
559                    # clean up in crud.py).
560                    # Since we have cols_to_return, just return what we asked
561                    # for (plus primary key, which ORM persistence needs since
562                    # we likely set bookkeeping=True here, which is another
563                    # whole thing...).   We dont want to clutter the
564                    # statement up with lots of other cols the user didn't
565                    # ask for.  see #9685
566                    *dml_mapper.primary_key,
567                    supplemental_cols=cols_to_return,
568                )
569            else:
570                dml_level_statement = dml_level_statement.returning(
571                    *cols_to_return
572                )
573
574        return dml_level_statement
575
576    @classmethod
577    def _return_orm_returning(
578        cls,
579        session,
580        statement,
581        params,
582        execution_options,
583        bind_arguments,
584        result,
585    ):
586        execution_context = result.context
587        compile_state = execution_context.compiled.compile_state
588
589        if (
590            compile_state.from_statement_ctx
591            and not compile_state.from_statement_ctx.compile_options._is_star
592        ):
593            load_options = execution_options.get(
594                "_sa_orm_load_options", QueryContext.default_load_options
595            )
596
597            querycontext = QueryContext(
598                compile_state.from_statement_ctx,
599                compile_state.select_statement,
600                params,
601                session,
602                load_options,
603                execution_options,
604                bind_arguments,
605            )
606            return loading.instances(result, querycontext)
607        else:
608            return result
609
610
611class BulkUDCompileState(ORMDMLState):
612    class default_update_options(Options):
613        _dml_strategy: DMLStrategyArgument = "auto"
614        _synchronize_session: SynchronizeSessionArgument = "auto"
615        _can_use_returning: bool = False
616        _is_delete_using: bool = False
617        _is_update_from: bool = False
618        _autoflush: bool = True
619        _subject_mapper: Optional[Mapper[Any]] = None
620        _resolved_values = EMPTY_DICT
621        _eval_condition = None
622        _matched_rows = None
623        _identity_token = None
624
625    @classmethod
626    def can_use_returning(
627        cls,
628        dialect: Dialect,
629        mapper: Mapper[Any],
630        *,
631        is_multitable: bool = False,
632        is_update_from: bool = False,
633        is_delete_using: bool = False,
634        is_executemany: bool = False,
635    ) -> bool:
636        raise NotImplementedError()
637
638    @classmethod
639    def orm_pre_session_exec(
640        cls,
641        session,
642        statement,
643        params,
644        execution_options,
645        bind_arguments,
646        is_pre_event,
647    ):
648        (
649            update_options,
650            execution_options,
651        ) = BulkUDCompileState.default_update_options.from_execution_options(
652            "_sa_orm_update_options",
653            {
654                "synchronize_session",
655                "autoflush",
656                "identity_token",
657                "is_delete_using",
658                "is_update_from",
659                "dml_strategy",
660            },
661            execution_options,
662            statement._execution_options,
663        )
664        bind_arguments["clause"] = statement
665        try:
666            plugin_subject = statement._propagate_attrs["plugin_subject"]
667        except KeyError:
668            assert False, "statement had 'orm' plugin but no plugin_subject"
669        else:
670            if plugin_subject:
671                bind_arguments["mapper"] = plugin_subject.mapper
672                update_options += {"_subject_mapper": plugin_subject.mapper}
673
674        if "parententity" not in statement.table._annotations:
675            update_options += {"_dml_strategy": "core_only"}
676        elif not isinstance(params, list):
677            if update_options._dml_strategy == "auto":
678                update_options += {"_dml_strategy": "orm"}
679            elif update_options._dml_strategy == "bulk":
680                raise sa_exc.InvalidRequestError(
681                    'Can\'t use "bulk" ORM insert strategy without '
682                    "passing separate parameters"
683                )
684        else:
685            if update_options._dml_strategy == "auto":
686                update_options += {"_dml_strategy": "bulk"}
687
688        sync = update_options._synchronize_session
689        if sync is not None:
690            if sync not in ("auto", "evaluate", "fetch", False):
691                raise sa_exc.ArgumentError(
692                    "Valid strategies for session synchronization "
693                    "are 'auto', 'evaluate', 'fetch', False"
694                )
695            if update_options._dml_strategy == "bulk" and sync == "fetch":
696                raise sa_exc.InvalidRequestError(
697                    "The 'fetch' synchronization strategy is not available "
698                    "for 'bulk' ORM updates (i.e. multiple parameter sets)"
699                )
700
701        if not is_pre_event:
702            if update_options._autoflush:
703                session._autoflush()
704
705            if update_options._dml_strategy == "orm":
706                if update_options._synchronize_session == "auto":
707                    update_options = cls._do_pre_synchronize_auto(
708                        session,
709                        statement,
710                        params,
711                        execution_options,
712                        bind_arguments,
713                        update_options,
714                    )
715                elif update_options._synchronize_session == "evaluate":
716                    update_options = cls._do_pre_synchronize_evaluate(
717                        session,
718                        statement,
719                        params,
720                        execution_options,
721                        bind_arguments,
722                        update_options,
723                    )
724                elif update_options._synchronize_session == "fetch":
725                    update_options = cls._do_pre_synchronize_fetch(
726                        session,
727                        statement,
728                        params,
729                        execution_options,
730                        bind_arguments,
731                        update_options,
732                    )
733            elif update_options._dml_strategy == "bulk":
734                if update_options._synchronize_session == "auto":
735                    update_options += {"_synchronize_session": "evaluate"}
736
737            # indicators from the "pre exec" step that are then
738            # added to the DML statement, which will also be part of the cache
739            # key.  The compile level create_for_statement() method will then
740            # consume these at compiler time.
741            statement = statement._annotate(
742                {
743                    "synchronize_session": update_options._synchronize_session,
744                    "is_delete_using": update_options._is_delete_using,
745                    "is_update_from": update_options._is_update_from,
746                    "dml_strategy": update_options._dml_strategy,
747                    "can_use_returning": update_options._can_use_returning,
748                }
749            )
750
751        return (
752            statement,
753            util.immutabledict(execution_options).union(
754                {"_sa_orm_update_options": update_options}
755            ),
756        )
757
758    @classmethod
759    def orm_setup_cursor_result(
760        cls,
761        session,
762        statement,
763        params,
764        execution_options,
765        bind_arguments,
766        result,
767    ):
768        # this stage of the execution is called after the
769        # do_orm_execute event hook.  meaning for an extension like
770        # horizontal sharding, this step happens *within* the horizontal
771        # sharding event handler which calls session.execute() re-entrantly
772        # and will occur for each backend individually.
773        # the sharding extension then returns its own merged result from the
774        # individual ones we return here.
775
776        update_options = execution_options["_sa_orm_update_options"]
777        if update_options._dml_strategy == "orm":
778            if update_options._synchronize_session == "evaluate":
779                cls._do_post_synchronize_evaluate(
780                    session, statement, result, update_options
781                )
782            elif update_options._synchronize_session == "fetch":
783                cls._do_post_synchronize_fetch(
784                    session, statement, result, update_options
785                )
786        elif update_options._dml_strategy == "bulk":
787            if update_options._synchronize_session == "evaluate":
788                cls._do_post_synchronize_bulk_evaluate(
789                    session, params, result, update_options
790                )
791            return result
792
793        return cls._return_orm_returning(
794            session,
795            statement,
796            params,
797            execution_options,
798            bind_arguments,
799            result,
800        )
801
802    @classmethod
803    def _adjust_for_extra_criteria(cls, global_attributes, ext_info):
804        """Apply extra criteria filtering.
805
806        For all distinct single-table-inheritance mappers represented in the
807        table being updated or deleted, produce additional WHERE criteria such
808        that only the appropriate subtypes are selected from the total results.
809
810        Additionally, add WHERE criteria originating from LoaderCriteriaOptions
811        collected from the statement.
812
813        """
814
815        return_crit = ()
816
817        adapter = ext_info._adapter if ext_info.is_aliased_class else None
818
819        if (
820            "additional_entity_criteria",
821            ext_info.mapper,
822        ) in global_attributes:
823            return_crit += tuple(
824                ae._resolve_where_criteria(ext_info)
825                for ae in global_attributes[
826                    ("additional_entity_criteria", ext_info.mapper)
827                ]
828                if ae.include_aliases or ae.entity is ext_info
829            )
830
831        if ext_info.mapper._single_table_criterion is not None:
832            return_crit += (ext_info.mapper._single_table_criterion,)
833
834        if adapter:
835            return_crit = tuple(adapter.traverse(crit) for crit in return_crit)
836
837        return return_crit
838
839    @classmethod
840    def _interpret_returning_rows(cls, mapper, rows):
841        """translate from local inherited table columns to base mapper
842        primary key columns.
843
844        Joined inheritance mappers always establish the primary key in terms of
845        the base table.   When we UPDATE a sub-table, we can only get
846        RETURNING for the sub-table's columns.
847
848        Here, we create a lookup from the local sub table's primary key
849        columns to the base table PK columns so that we can get identity
850        key values from RETURNING that's against the joined inheritance
851        sub-table.
852
853        the complexity here is to support more than one level deep of
854        inheritance, where we have to link columns to each other across
855        the inheritance hierarchy.
856
857        """
858
859        if mapper.local_table is not mapper.base_mapper.local_table:
860            return rows
861
862        # this starts as a mapping of
863        # local_pk_col: local_pk_col.
864        # we will then iteratively rewrite the "value" of the dict with
865        # each successive superclass column
866        local_pk_to_base_pk = {pk: pk for pk in mapper.local_table.primary_key}
867
868        for mp in mapper.iterate_to_root():
869            if mp.inherits is None:
870                break
871            elif mp.local_table is mp.inherits.local_table:
872                continue
873
874            t_to_e = dict(mp._table_to_equated[mp.inherits.local_table])
875            col_to_col = {sub_pk: super_pk for super_pk, sub_pk in t_to_e[mp]}
876            for pk, super_ in local_pk_to_base_pk.items():
877                local_pk_to_base_pk[pk] = col_to_col[super_]
878
879        lookup = {
880            local_pk_to_base_pk[lpk]: idx
881            for idx, lpk in enumerate(mapper.local_table.primary_key)
882        }
883        primary_key_convert = [
884            lookup[bpk] for bpk in mapper.base_mapper.primary_key
885        ]
886        return [tuple(row[idx] for idx in primary_key_convert) for row in rows]
887
888    @classmethod
889    def _get_matched_objects_on_criteria(cls, update_options, states):
890        mapper = update_options._subject_mapper
891        eval_condition = update_options._eval_condition
892
893        raw_data = [
894            (state.obj(), state, state.dict)
895            for state in states
896            if state.mapper.isa(mapper) and not state.expired
897        ]
898
899        identity_token = update_options._identity_token
900        if identity_token is not None:
901            raw_data = [
902                (obj, state, dict_)
903                for obj, state, dict_ in raw_data
904                if state.identity_token == identity_token
905            ]
906
907        result = []
908        for obj, state, dict_ in raw_data:
909            evaled_condition = eval_condition(obj)
910
911            # caution: don't use "in ()" or == here, _EXPIRE_OBJECT
912            # evaluates as True for all comparisons
913            if (
914                evaled_condition is True
915                or evaled_condition is evaluator._EXPIRED_OBJECT
916            ):
917                result.append(
918                    (
919                        obj,
920                        state,
921                        dict_,
922                        evaled_condition is evaluator._EXPIRED_OBJECT,
923                    )
924                )
925        return result
926
927    @classmethod
928    def _eval_condition_from_statement(cls, update_options, statement):
929        mapper = update_options._subject_mapper
930        target_cls = mapper.class_
931
932        evaluator_compiler = evaluator._EvaluatorCompiler(target_cls)
933        crit = ()
934        if statement._where_criteria:
935            crit += statement._where_criteria
936
937        global_attributes = {}
938        for opt in statement._with_options:
939            if opt._is_criteria_option:
940                opt.get_global_criteria(global_attributes)
941
942        if global_attributes:
943            crit += cls._adjust_for_extra_criteria(global_attributes, mapper)
944
945        if crit:
946            eval_condition = evaluator_compiler.process(*crit)
947        else:
948            # workaround for mypy https://github.com/python/mypy/issues/14027
949            def _eval_condition(obj):
950                return True
951
952            eval_condition = _eval_condition
953
954        return eval_condition
955
956    @classmethod
957    def _do_pre_synchronize_auto(
958        cls,
959        session,
960        statement,
961        params,
962        execution_options,
963        bind_arguments,
964        update_options,
965    ):
966        """setup auto sync strategy
967
968
969        "auto" checks if we can use "evaluate" first, then falls back
970        to "fetch"
971
972        evaluate is vastly more efficient for the common case
973        where session is empty, only has a few objects, and the UPDATE
974        statement can potentially match thousands/millions of rows.
975
976        OTOH more complex criteria that fails to work with "evaluate"
977        we would hope usually correlates with fewer net rows.
978
979        """
980
981        try:
982            eval_condition = cls._eval_condition_from_statement(
983                update_options, statement
984            )
985
986        except evaluator.UnevaluatableError:
987            pass
988        else:
989            return update_options + {
990                "_eval_condition": eval_condition,
991                "_synchronize_session": "evaluate",
992            }
993
994        update_options += {"_synchronize_session": "fetch"}
995        return cls._do_pre_synchronize_fetch(
996            session,
997            statement,
998            params,
999            execution_options,
1000            bind_arguments,
1001            update_options,
1002        )
1003
1004    @classmethod
1005    def _do_pre_synchronize_evaluate(
1006        cls,
1007        session,
1008        statement,
1009        params,
1010        execution_options,
1011        bind_arguments,
1012        update_options,
1013    ):
1014        try:
1015            eval_condition = cls._eval_condition_from_statement(
1016                update_options, statement
1017            )
1018
1019        except evaluator.UnevaluatableError as err:
1020            raise sa_exc.InvalidRequestError(
1021                'Could not evaluate current criteria in Python: "%s". '
1022                "Specify 'fetch' or False for the "
1023                "synchronize_session execution option." % err
1024            ) from err
1025
1026        return update_options + {
1027            "_eval_condition": eval_condition,
1028        }
1029
1030    @classmethod
1031    def _get_resolved_values(cls, mapper, statement):
1032        if statement._multi_values:
1033            return []
1034        elif statement._ordered_values:
1035            return list(statement._ordered_values)
1036        elif statement._values:
1037            return list(statement._values.items())
1038        else:
1039            return []
1040
1041    @classmethod
1042    def _resolved_keys_as_propnames(cls, mapper, resolved_values):
1043        values = []
1044        for k, v in resolved_values:
1045            if mapper and isinstance(k, expression.ColumnElement):
1046                try:
1047                    attr = mapper._columntoproperty[k]
1048                except orm_exc.UnmappedColumnError:
1049                    pass
1050                else:
1051                    values.append((attr.key, v))
1052            else:
1053                raise sa_exc.InvalidRequestError(
1054                    "Attribute name not found, can't be "
1055                    "synchronized back to objects: %r" % k
1056                )
1057        return values
1058
1059    @classmethod
1060    def _do_pre_synchronize_fetch(
1061        cls,
1062        session,
1063        statement,
1064        params,
1065        execution_options,
1066        bind_arguments,
1067        update_options,
1068    ):
1069        mapper = update_options._subject_mapper
1070
1071        select_stmt = (
1072            select(*(mapper.primary_key + (mapper.select_identity_token,)))
1073            .select_from(mapper)
1074            .options(*statement._with_options)
1075        )
1076        select_stmt._where_criteria = statement._where_criteria
1077
1078        # conditionally run the SELECT statement for pre-fetch, testing the
1079        # "bind" for if we can use RETURNING or not using the do_orm_execute
1080        # event.  If RETURNING is available, the do_orm_execute event
1081        # will cancel the SELECT from being actually run.
1082        #
1083        # The way this is organized seems strange, why don't we just
1084        # call can_use_returning() before invoking the statement and get
1085        # answer?, why does this go through the whole execute phase using an
1086        # event?  Answer: because we are integrating with extensions such
1087        # as the horizontal sharding extention that "multiplexes" an individual
1088        # statement run through multiple engines, and it uses
1089        # do_orm_execute() to do that.
1090
1091        can_use_returning = None
1092
1093        def skip_for_returning(orm_context: ORMExecuteState) -> Any:
1094            bind = orm_context.session.get_bind(**orm_context.bind_arguments)
1095            nonlocal can_use_returning
1096
1097            per_bind_result = cls.can_use_returning(
1098                bind.dialect,
1099                mapper,
1100                is_update_from=update_options._is_update_from,
1101                is_delete_using=update_options._is_delete_using,
1102                is_executemany=orm_context.is_executemany,
1103            )
1104
1105            if can_use_returning is not None:
1106                if can_use_returning != per_bind_result:
1107                    raise sa_exc.InvalidRequestError(
1108                        "For synchronize_session='fetch', can't mix multiple "
1109                        "backends where some support RETURNING and others "
1110                        "don't"
1111                    )
1112            elif orm_context.is_executemany and not per_bind_result:
1113                raise sa_exc.InvalidRequestError(
1114                    "For synchronize_session='fetch', can't use multiple "
1115                    "parameter sets in ORM mode, which this backend does not "
1116                    "support with RETURNING"
1117                )
1118            else:
1119                can_use_returning = per_bind_result
1120
1121            if per_bind_result:
1122                return _result.null_result()
1123            else:
1124                return None
1125
1126        result = session.execute(
1127            select_stmt,
1128            params,
1129            execution_options=execution_options,
1130            bind_arguments=bind_arguments,
1131            _add_event=skip_for_returning,
1132        )
1133        matched_rows = result.fetchall()
1134
1135        return update_options + {
1136            "_matched_rows": matched_rows,
1137            "_can_use_returning": can_use_returning,
1138        }
1139
1140
1141@CompileState.plugin_for("orm", "insert")
1142class BulkORMInsert(ORMDMLState, InsertDMLState):
1143    class default_insert_options(Options):
1144        _dml_strategy: DMLStrategyArgument = "auto"
1145        _render_nulls: bool = False
1146        _return_defaults: bool = False
1147        _subject_mapper: Optional[Mapper[Any]] = None
1148        _autoflush: bool = True
1149        _populate_existing: bool = False
1150
1151    select_statement: Optional[FromStatement] = None
1152
1153    @classmethod
1154    def orm_pre_session_exec(
1155        cls,
1156        session,
1157        statement,
1158        params,
1159        execution_options,
1160        bind_arguments,
1161        is_pre_event,
1162    ):
1163        (
1164            insert_options,
1165            execution_options,
1166        ) = BulkORMInsert.default_insert_options.from_execution_options(
1167            "_sa_orm_insert_options",
1168            {"dml_strategy", "autoflush", "populate_existing", "render_nulls"},
1169            execution_options,
1170            statement._execution_options,
1171        )
1172        bind_arguments["clause"] = statement
1173        try:
1174            plugin_subject = statement._propagate_attrs["plugin_subject"]
1175        except KeyError:
1176            assert False, "statement had 'orm' plugin but no plugin_subject"
1177        else:
1178            if plugin_subject:
1179                bind_arguments["mapper"] = plugin_subject.mapper
1180                insert_options += {"_subject_mapper": plugin_subject.mapper}
1181
1182        if not params:
1183            if insert_options._dml_strategy == "auto":
1184                insert_options += {"_dml_strategy": "orm"}
1185            elif insert_options._dml_strategy == "bulk":
1186                raise sa_exc.InvalidRequestError(
1187                    'Can\'t use "bulk" ORM insert strategy without '
1188                    "passing separate parameters"
1189                )
1190        else:
1191            if insert_options._dml_strategy == "auto":
1192                insert_options += {"_dml_strategy": "bulk"}
1193
1194        if insert_options._dml_strategy != "raw":
1195            # for ORM object loading, like ORMContext, we have to disable
1196            # result set adapt_to_context, because we will be generating a
1197            # new statement with specific columns that's cached inside of
1198            # an ORMFromStatementCompileState, which we will re-use for
1199            # each result.
1200            if not execution_options:

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

codekingpro/portable-devtools · Team Ai