codekingpro/portable-devtools
114k
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:
