Remove Elixir library
Update SQLAlchemy
This commit is contained in:
@@ -1,5 +1,5 @@
|
||||
# orm/persistence.py
|
||||
# Copyright (C) 2005-2013 the SQLAlchemy authors and contributors <see AUTHORS file>
|
||||
# Copyright (C) 2005-2014 the SQLAlchemy authors and contributors <see AUTHORS file>
|
||||
#
|
||||
# This module is part of SQLAlchemy and is released under
|
||||
# the MIT License: http://www.opensource.org/licenses/mit-license.php
|
||||
@@ -15,12 +15,12 @@ in unitofwork.py.
|
||||
|
||||
import operator
|
||||
from itertools import groupby
|
||||
from .. import sql, util, exc as sa_exc, schema
|
||||
from . import attributes, sync, exc as orm_exc, evaluator
|
||||
from .base import _state_mapper, state_str, _attr_as_key
|
||||
from ..sql import expression
|
||||
from . import loading
|
||||
|
||||
from sqlalchemy import sql, util, exc as sa_exc
|
||||
from sqlalchemy.orm import attributes, sync, \
|
||||
exc as orm_exc
|
||||
|
||||
from sqlalchemy.orm.util import _state_mapper, state_str
|
||||
|
||||
def save_obj(base_mapper, states, uowtransaction, single=False):
|
||||
"""Issue ``INSERT`` and/or ``UPDATE`` statements for a list
|
||||
@@ -46,7 +46,7 @@ def save_obj(base_mapper, states, uowtransaction, single=False):
|
||||
|
||||
cached_connections = _cached_connection_dict(base_mapper)
|
||||
|
||||
for table, mapper in base_mapper._sorted_tables.iteritems():
|
||||
for table, mapper in base_mapper._sorted_tables.items():
|
||||
insert = _collect_insert_commands(base_mapper, uowtransaction,
|
||||
table, states_to_insert)
|
||||
|
||||
@@ -61,11 +61,12 @@ def save_obj(base_mapper, states, uowtransaction, single=False):
|
||||
if insert:
|
||||
_emit_insert_statements(base_mapper, uowtransaction,
|
||||
cached_connections,
|
||||
table, insert)
|
||||
mapper, table, insert)
|
||||
|
||||
_finalize_insert_update_commands(base_mapper, uowtransaction,
|
||||
states_to_insert, states_to_update)
|
||||
|
||||
|
||||
def post_update(base_mapper, states, uowtransaction, post_update_cols):
|
||||
"""Issue UPDATE statements on behalf of a relationship() which
|
||||
specifies post_update.
|
||||
@@ -77,8 +78,7 @@ def post_update(base_mapper, states, uowtransaction, post_update_cols):
|
||||
base_mapper,
|
||||
states, uowtransaction)
|
||||
|
||||
|
||||
for table, mapper in base_mapper._sorted_tables.iteritems():
|
||||
for table, mapper in base_mapper._sorted_tables.items():
|
||||
update = _collect_post_update_commands(base_mapper, uowtransaction,
|
||||
table, states_to_update,
|
||||
post_update_cols)
|
||||
@@ -88,6 +88,7 @@ def post_update(base_mapper, states, uowtransaction, post_update_cols):
|
||||
cached_connections,
|
||||
mapper, table, update)
|
||||
|
||||
|
||||
def delete_obj(base_mapper, states, uowtransaction):
|
||||
"""Issue ``DELETE`` statements for a list of objects.
|
||||
|
||||
@@ -105,7 +106,7 @@ def delete_obj(base_mapper, states, uowtransaction):
|
||||
|
||||
table_to_mapper = base_mapper._sorted_tables
|
||||
|
||||
for table in reversed(table_to_mapper.keys()):
|
||||
for table in reversed(list(table_to_mapper.keys())):
|
||||
delete = _collect_delete_commands(base_mapper, uowtransaction,
|
||||
table, states_to_delete)
|
||||
|
||||
@@ -118,6 +119,7 @@ def delete_obj(base_mapper, states, uowtransaction):
|
||||
in states_to_delete:
|
||||
mapper.dispatch.after_delete(mapper, connection, state)
|
||||
|
||||
|
||||
def _organize_states_for_save(base_mapper, states, uowtransaction):
|
||||
"""Make an initial pass across a set of states for INSERT or
|
||||
UPDATE.
|
||||
@@ -148,12 +150,15 @@ def _organize_states_for_save(base_mapper, states, uowtransaction):
|
||||
else:
|
||||
mapper.dispatch.before_update(mapper, connection, state)
|
||||
|
||||
if mapper._validate_polymorphic_identity:
|
||||
mapper._validate_polymorphic_identity(mapper, state, dict_)
|
||||
|
||||
# detect if we have a "pending" instance (i.e. has
|
||||
# no instance_key attached to it), and another instance
|
||||
# with the same identity key already exists as persistent.
|
||||
# convert to an UPDATE if so.
|
||||
if not has_identity and \
|
||||
instance_key in uowtransaction.session.identity_map:
|
||||
instance_key in uowtransaction.session.identity_map:
|
||||
instance = \
|
||||
uowtransaction.session.identity_map[instance_key]
|
||||
existing = attributes.instance_state(instance)
|
||||
@@ -187,6 +192,7 @@ def _organize_states_for_save(base_mapper, states, uowtransaction):
|
||||
|
||||
return states_to_insert, states_to_update
|
||||
|
||||
|
||||
def _organize_states_for_post_update(base_mapper, states,
|
||||
uowtransaction):
|
||||
"""Make an initial pass across a set of states for UPDATE
|
||||
@@ -200,6 +206,7 @@ def _organize_states_for_post_update(base_mapper, states,
|
||||
return list(_connections_for_states(base_mapper, uowtransaction,
|
||||
states))
|
||||
|
||||
|
||||
def _organize_states_for_delete(base_mapper, states, uowtransaction):
|
||||
"""Make an initial pass across a set of states for DELETE.
|
||||
|
||||
@@ -220,6 +227,7 @@ def _organize_states_for_delete(base_mapper, states, uowtransaction):
|
||||
bool(state.key), connection))
|
||||
return states_to_delete
|
||||
|
||||
|
||||
def _collect_insert_commands(base_mapper, uowtransaction, table,
|
||||
states_to_insert):
|
||||
"""Identify sets of values to use in INSERT statements for a
|
||||
@@ -238,9 +246,12 @@ def _collect_insert_commands(base_mapper, uowtransaction, table,
|
||||
value_params = {}
|
||||
|
||||
has_all_pks = True
|
||||
has_all_defaults = True
|
||||
for col in mapper._cols_by_table[table]:
|
||||
if col is mapper.version_id_col:
|
||||
params[col.key] = mapper.version_id_generator(None)
|
||||
if col is mapper.version_id_col and \
|
||||
mapper.version_id_generator is not False:
|
||||
val = mapper.version_id_generator(None)
|
||||
params[col.key] = val
|
||||
else:
|
||||
# pull straight from the dict for
|
||||
# pending objects
|
||||
@@ -253,6 +264,9 @@ def _collect_insert_commands(base_mapper, uowtransaction, table,
|
||||
elif col.default is None and \
|
||||
col.server_default is None:
|
||||
params[col.key] = value
|
||||
elif col.server_default is not None and \
|
||||
mapper.base_mapper.eager_defaults:
|
||||
has_all_defaults = False
|
||||
|
||||
elif isinstance(value, sql.ClauseElement):
|
||||
value_params[col] = value
|
||||
@@ -260,9 +274,11 @@ def _collect_insert_commands(base_mapper, uowtransaction, table,
|
||||
params[col.key] = value
|
||||
|
||||
insert.append((state, state_dict, params, mapper,
|
||||
connection, value_params, has_all_pks))
|
||||
connection, value_params, has_all_pks,
|
||||
has_all_defaults))
|
||||
return insert
|
||||
|
||||
|
||||
def _collect_update_commands(base_mapper, uowtransaction,
|
||||
table, states_to_update):
|
||||
"""Identify sets of values to use in UPDATE statements for a
|
||||
@@ -306,19 +322,20 @@ def _collect_update_commands(base_mapper, uowtransaction,
|
||||
params[col.key] = history.added[0]
|
||||
hasdata = True
|
||||
else:
|
||||
params[col.key] = mapper.version_id_generator(
|
||||
params[col._label])
|
||||
if mapper.version_id_generator is not False:
|
||||
val = mapper.version_id_generator(params[col._label])
|
||||
params[col.key] = val
|
||||
|
||||
# HACK: check for history, in case the
|
||||
# history is only
|
||||
# in a different table than the one
|
||||
# where the version_id_col is.
|
||||
for prop in mapper._columntoproperty.itervalues():
|
||||
history = attributes.get_state_history(
|
||||
state, prop.key,
|
||||
attributes.PASSIVE_NO_INITIALIZE)
|
||||
if history.added:
|
||||
hasdata = True
|
||||
# HACK: check for history, in case the
|
||||
# history is only
|
||||
# in a different table than the one
|
||||
# where the version_id_col is.
|
||||
for prop in mapper._columntoproperty.values():
|
||||
history = attributes.get_state_history(
|
||||
state, prop.key,
|
||||
attributes.PASSIVE_NO_INITIALIZE)
|
||||
if history.added:
|
||||
hasdata = True
|
||||
else:
|
||||
prop = mapper._columntoproperty[col]
|
||||
history = attributes.get_state_history(
|
||||
@@ -370,7 +387,7 @@ def _collect_update_commands(base_mapper, uowtransaction,
|
||||
params[col._label] = value
|
||||
if hasdata:
|
||||
if hasnull:
|
||||
raise sa_exc.FlushError(
|
||||
raise orm_exc.FlushError(
|
||||
"Can't update table "
|
||||
"using NULL for primary "
|
||||
"key value")
|
||||
@@ -400,6 +417,7 @@ def _collect_post_update_commands(base_mapper, uowtransaction, table,
|
||||
mapper._get_state_attr_by_column(
|
||||
state,
|
||||
state_dict, col)
|
||||
|
||||
elif col in post_update_cols:
|
||||
prop = mapper._columntoproperty[col]
|
||||
history = attributes.get_state_history(
|
||||
@@ -414,6 +432,7 @@ def _collect_post_update_commands(base_mapper, uowtransaction, table,
|
||||
connection))
|
||||
return update
|
||||
|
||||
|
||||
def _collect_delete_commands(base_mapper, uowtransaction, table,
|
||||
states_to_delete):
|
||||
"""Identify values to use in DELETE statements for a list of
|
||||
@@ -434,7 +453,7 @@ def _collect_delete_commands(base_mapper, uowtransaction, table,
|
||||
mapper._get_state_attr_by_column(
|
||||
state, state_dict, col)
|
||||
if value is None:
|
||||
raise sa_exc.FlushError(
|
||||
raise orm_exc.FlushError(
|
||||
"Can't delete from table "
|
||||
"using NULL for primary "
|
||||
"key value")
|
||||
@@ -468,7 +487,13 @@ def _emit_update_statements(base_mapper, uowtransaction,
|
||||
sql.bindparam(mapper.version_id_col._label,
|
||||
type_=mapper.version_id_col.type))
|
||||
|
||||
return table.update(clause)
|
||||
stmt = table.update(clause)
|
||||
if mapper.base_mapper.eager_defaults:
|
||||
stmt = stmt.return_defaults()
|
||||
elif mapper.version_id_col is not None:
|
||||
stmt = stmt.return_defaults(mapper.version_id_col)
|
||||
|
||||
return stmt
|
||||
|
||||
statement = base_mapper._memo(('update', table), update_stmt)
|
||||
|
||||
@@ -490,8 +515,7 @@ def _emit_update_statements(base_mapper, uowtransaction,
|
||||
table,
|
||||
state,
|
||||
state_dict,
|
||||
c.context.prefetch_cols,
|
||||
c.context.postfetch_cols,
|
||||
c,
|
||||
c.context.compiled_parameters[0],
|
||||
value_params)
|
||||
rows += c.rowcount
|
||||
@@ -509,45 +533,57 @@ def _emit_update_statements(base_mapper, uowtransaction,
|
||||
c.dialect.dialect_description,
|
||||
stacklevel=12)
|
||||
|
||||
|
||||
def _emit_insert_statements(base_mapper, uowtransaction,
|
||||
cached_connections, table, insert):
|
||||
cached_connections, mapper, table, insert):
|
||||
"""Emit INSERT statements corresponding to value lists collected
|
||||
by _collect_insert_commands()."""
|
||||
|
||||
statement = base_mapper._memo(('insert', table), table.insert)
|
||||
|
||||
for (connection, pkeys, hasvalue, has_all_pks), \
|
||||
for (connection, pkeys, hasvalue, has_all_pks, has_all_defaults), \
|
||||
records in groupby(insert,
|
||||
lambda rec: (rec[4],
|
||||
rec[2].keys(),
|
||||
list(rec[2].keys()),
|
||||
bool(rec[5]),
|
||||
rec[6])
|
||||
rec[6], rec[7])
|
||||
):
|
||||
if has_all_pks and not hasvalue:
|
||||
if \
|
||||
(
|
||||
has_all_defaults
|
||||
or not base_mapper.eager_defaults
|
||||
or not connection.dialect.implicit_returning
|
||||
) and has_all_pks and not hasvalue:
|
||||
|
||||
records = list(records)
|
||||
multiparams = [rec[2] for rec in records]
|
||||
|
||||
c = cached_connections[connection].\
|
||||
execute(statement, multiparams)
|
||||
|
||||
for (state, state_dict, params, mapper,
|
||||
conn, value_params, has_all_pks), \
|
||||
for (state, state_dict, params, mapper_rec,
|
||||
conn, value_params, has_all_pks, has_all_defaults), \
|
||||
last_inserted_params in \
|
||||
zip(records, c.context.compiled_parameters):
|
||||
_postfetch(
|
||||
mapper,
|
||||
mapper_rec,
|
||||
uowtransaction,
|
||||
table,
|
||||
state,
|
||||
state_dict,
|
||||
c.context.prefetch_cols,
|
||||
c.context.postfetch_cols,
|
||||
c,
|
||||
last_inserted_params,
|
||||
value_params)
|
||||
|
||||
else:
|
||||
for state, state_dict, params, mapper, \
|
||||
if not has_all_defaults and base_mapper.eager_defaults:
|
||||
statement = statement.return_defaults()
|
||||
elif mapper.version_id_col is not None:
|
||||
statement = statement.return_defaults(mapper.version_id_col)
|
||||
|
||||
for state, state_dict, params, mapper_rec, \
|
||||
connection, value_params, \
|
||||
has_all_pks in records:
|
||||
has_all_pks, has_all_defaults in records:
|
||||
|
||||
if value_params:
|
||||
result = connection.execute(
|
||||
@@ -563,28 +599,26 @@ def _emit_insert_statements(base_mapper, uowtransaction,
|
||||
# set primary key attributes
|
||||
for pk, col in zip(primary_key,
|
||||
mapper._pks_by_table[table]):
|
||||
prop = mapper._columntoproperty[col]
|
||||
prop = mapper_rec._columntoproperty[col]
|
||||
if state_dict.get(prop.key) is None:
|
||||
# TODO: would rather say:
|
||||
#state_dict[prop.key] = pk
|
||||
mapper._set_state_attr_by_column(
|
||||
mapper_rec._set_state_attr_by_column(
|
||||
state,
|
||||
state_dict,
|
||||
col, pk)
|
||||
|
||||
_postfetch(
|
||||
mapper,
|
||||
mapper_rec,
|
||||
uowtransaction,
|
||||
table,
|
||||
state,
|
||||
state_dict,
|
||||
result.context.prefetch_cols,
|
||||
result.context.postfetch_cols,
|
||||
result,
|
||||
result.context.compiled_parameters[0],
|
||||
value_params)
|
||||
|
||||
|
||||
|
||||
def _emit_post_update_statements(base_mapper, uowtransaction,
|
||||
cached_connections, mapper, table, update):
|
||||
"""Emit UPDATE statements corresponding to value lists collected
|
||||
@@ -606,7 +640,7 @@ def _emit_post_update_statements(base_mapper, uowtransaction,
|
||||
# also group them into common (connection, cols) sets
|
||||
# to support executemany().
|
||||
for key, grouper in groupby(
|
||||
update, lambda rec: (rec[4], rec[2].keys())
|
||||
update, lambda rec: (rec[4], list(rec[2].keys()))
|
||||
):
|
||||
connection = key[0]
|
||||
multiparams = [params for state, state_dict,
|
||||
@@ -640,7 +674,7 @@ def _emit_delete_statements(base_mapper, uowtransaction, cached_connections,
|
||||
|
||||
return table.delete(clause)
|
||||
|
||||
for connection, del_objects in delete.iteritems():
|
||||
for connection, del_objects in delete.items():
|
||||
statement = base_mapper._memo(('delete', table), delete_stmt)
|
||||
|
||||
connection = cached_connections[connection]
|
||||
@@ -687,15 +721,27 @@ def _finalize_insert_update_commands(base_mapper, uowtransaction,
|
||||
if p.expire_on_flush or p.key not in state.dict]
|
||||
)
|
||||
if readonly:
|
||||
state.expire_attributes(state.dict, readonly)
|
||||
state._expire_attributes(state.dict, readonly)
|
||||
|
||||
# if eager_defaults option is enabled,
|
||||
# refresh whatever has been expired.
|
||||
if base_mapper.eager_defaults and state.unloaded:
|
||||
# if eager_defaults option is enabled, load
|
||||
# all expired cols. Else if we have a version_id_col, make sure
|
||||
# it isn't expired.
|
||||
toload_now = []
|
||||
|
||||
if base_mapper.eager_defaults:
|
||||
toload_now.extend(state._unloaded_non_object)
|
||||
elif mapper.version_id_col is not None and \
|
||||
mapper.version_id_generator is False:
|
||||
prop = mapper._columntoproperty[mapper.version_id_col]
|
||||
if prop.key in state.unloaded:
|
||||
toload_now.extend([prop.key])
|
||||
|
||||
if toload_now:
|
||||
state.key = base_mapper._identity_key_from_state(state)
|
||||
uowtransaction.session.query(base_mapper)._load_on_ident(
|
||||
loading.load_on_ident(
|
||||
uowtransaction.session.query(base_mapper),
|
||||
state.key, refresh_state=state,
|
||||
only_load_props=state.unloaded)
|
||||
only_load_props=toload_now)
|
||||
|
||||
# call after_XXX extensions
|
||||
if not has_identity:
|
||||
@@ -703,22 +749,34 @@ def _finalize_insert_update_commands(base_mapper, uowtransaction,
|
||||
else:
|
||||
mapper.dispatch.after_update(mapper, connection, state)
|
||||
|
||||
|
||||
def _postfetch(mapper, uowtransaction, table,
|
||||
state, dict_, prefetch_cols, postfetch_cols,
|
||||
params, value_params):
|
||||
state, dict_, result, params, value_params):
|
||||
"""Expire attributes in need of newly persisted database state,
|
||||
after an INSERT or UPDATE statement has proceeded for that
|
||||
state."""
|
||||
|
||||
prefetch_cols = result.context.prefetch_cols
|
||||
postfetch_cols = result.context.postfetch_cols
|
||||
returning_cols = result.context.returning_cols
|
||||
|
||||
if mapper.version_id_col is not None:
|
||||
prefetch_cols = list(prefetch_cols) + [mapper.version_id_col]
|
||||
|
||||
if returning_cols:
|
||||
row = result.context.returned_defaults
|
||||
if row is not None:
|
||||
for col in returning_cols:
|
||||
if col.primary_key:
|
||||
continue
|
||||
mapper._set_state_attr_by_column(state, dict_, col, row[col])
|
||||
|
||||
for c in prefetch_cols:
|
||||
if c.key in params and c in mapper._columntoproperty:
|
||||
mapper._set_state_attr_by_column(state, dict_, c, params[c.key])
|
||||
|
||||
if postfetch_cols:
|
||||
state.expire_attributes(state.dict,
|
||||
state._expire_attributes(state.dict,
|
||||
[mapper._columntoproperty[c].key
|
||||
for c in postfetch_cols if c in
|
||||
mapper._columntoproperty]
|
||||
@@ -733,6 +791,7 @@ def _postfetch(mapper, uowtransaction, table,
|
||||
uowtransaction,
|
||||
mapper.passive_updates)
|
||||
|
||||
|
||||
def _connections_for_states(base_mapper, uowtransaction, states):
|
||||
"""Return an iterator of (state, state.dict, mapper, connection).
|
||||
|
||||
@@ -762,18 +821,265 @@ def _connections_for_states(base_mapper, uowtransaction, states):
|
||||
|
||||
yield state, state.dict, mapper, connection
|
||||
|
||||
|
||||
def _cached_connection_dict(base_mapper):
|
||||
# dictionary of connection->connection_with_cache_options.
|
||||
return util.PopulateDict(
|
||||
lambda conn:conn.execution_options(
|
||||
lambda conn: conn.execution_options(
|
||||
compiled_cache=base_mapper._compiled_cache
|
||||
))
|
||||
|
||||
|
||||
def _sort_states(states):
|
||||
pending = set(states)
|
||||
persistent = set(s for s in pending if s.key is not None)
|
||||
pending.difference_update(persistent)
|
||||
return sorted(pending, key=operator.attrgetter("insert_order")) + \
|
||||
sorted(persistent, key=lambda q:q.key[1])
|
||||
sorted(persistent, key=lambda q: q.key[1])
|
||||
|
||||
|
||||
class BulkUD(object):
|
||||
"""Handle bulk update and deletes via a :class:`.Query`."""
|
||||
|
||||
def __init__(self, query):
|
||||
self.query = query.enable_eagerloads(False)
|
||||
|
||||
@property
|
||||
def session(self):
|
||||
return self.query.session
|
||||
|
||||
@classmethod
|
||||
def _factory(cls, lookup, synchronize_session, *arg):
|
||||
try:
|
||||
klass = lookup[synchronize_session]
|
||||
except KeyError:
|
||||
raise sa_exc.ArgumentError(
|
||||
"Valid strategies for session synchronization "
|
||||
"are %s" % (", ".join(sorted(repr(x)
|
||||
for x in lookup))))
|
||||
else:
|
||||
return klass(*arg)
|
||||
|
||||
def exec_(self):
|
||||
self._do_pre()
|
||||
self._do_pre_synchronize()
|
||||
self._do_exec()
|
||||
self._do_post_synchronize()
|
||||
self._do_post()
|
||||
|
||||
def _do_pre(self):
|
||||
query = self.query
|
||||
self.context = context = query._compile_context()
|
||||
if len(context.statement.froms) != 1 or \
|
||||
not isinstance(context.statement.froms[0], schema.Table):
|
||||
|
||||
self.primary_table = query._only_entity_zero(
|
||||
"This operation requires only one Table or "
|
||||
"entity be specified as the target."
|
||||
).mapper.local_table
|
||||
else:
|
||||
self.primary_table = context.statement.froms[0]
|
||||
|
||||
session = query.session
|
||||
|
||||
if query._autoflush:
|
||||
session._autoflush()
|
||||
|
||||
def _do_pre_synchronize(self):
|
||||
pass
|
||||
|
||||
def _do_post_synchronize(self):
|
||||
pass
|
||||
|
||||
|
||||
class BulkEvaluate(BulkUD):
|
||||
"""BulkUD which does the 'evaluate' method of session state resolution."""
|
||||
|
||||
def _additional_evaluators(self, evaluator_compiler):
|
||||
pass
|
||||
|
||||
def _do_pre_synchronize(self):
|
||||
query = self.query
|
||||
try:
|
||||
evaluator_compiler = evaluator.EvaluatorCompiler()
|
||||
if query.whereclause is not None:
|
||||
eval_condition = evaluator_compiler.process(
|
||||
query.whereclause)
|
||||
else:
|
||||
def eval_condition(obj):
|
||||
return True
|
||||
|
||||
self._additional_evaluators(evaluator_compiler)
|
||||
|
||||
except evaluator.UnevaluatableError:
|
||||
raise sa_exc.InvalidRequestError(
|
||||
"Could not evaluate current criteria in Python. "
|
||||
"Specify 'fetch' or False for the "
|
||||
"synchronize_session parameter.")
|
||||
target_cls = query._mapper_zero().class_
|
||||
|
||||
#TODO: detect when the where clause is a trivial primary key match
|
||||
self.matched_objects = [
|
||||
obj for (cls, pk), obj in
|
||||
query.session.identity_map.items()
|
||||
if issubclass(cls, target_cls) and
|
||||
eval_condition(obj)]
|
||||
|
||||
|
||||
class BulkFetch(BulkUD):
|
||||
"""BulkUD which does the 'fetch' method of session state resolution."""
|
||||
|
||||
def _do_pre_synchronize(self):
|
||||
query = self.query
|
||||
session = query.session
|
||||
select_stmt = self.context.statement.with_only_columns(
|
||||
self.primary_table.primary_key)
|
||||
self.matched_rows = session.execute(
|
||||
select_stmt,
|
||||
params=query._params).fetchall()
|
||||
|
||||
|
||||
class BulkUpdate(BulkUD):
|
||||
"""BulkUD which handles UPDATEs."""
|
||||
|
||||
def __init__(self, query, values):
|
||||
super(BulkUpdate, self).__init__(query)
|
||||
self.query._no_select_modifiers("update")
|
||||
self.values = values
|
||||
|
||||
@classmethod
|
||||
def factory(cls, query, synchronize_session, values):
|
||||
return BulkUD._factory({
|
||||
"evaluate": BulkUpdateEvaluate,
|
||||
"fetch": BulkUpdateFetch,
|
||||
False: BulkUpdate
|
||||
}, synchronize_session, query, values)
|
||||
|
||||
def _do_exec(self):
|
||||
update_stmt = sql.update(self.primary_table,
|
||||
self.context.whereclause, self.values)
|
||||
|
||||
self.result = self.query.session.execute(
|
||||
update_stmt, params=self.query._params)
|
||||
self.rowcount = self.result.rowcount
|
||||
|
||||
def _do_post(self):
|
||||
session = self.query.session
|
||||
session.dispatch.after_bulk_update(self)
|
||||
|
||||
|
||||
class BulkDelete(BulkUD):
|
||||
"""BulkUD which handles DELETEs."""
|
||||
|
||||
def __init__(self, query):
|
||||
super(BulkDelete, self).__init__(query)
|
||||
self.query._no_select_modifiers("delete")
|
||||
|
||||
@classmethod
|
||||
def factory(cls, query, synchronize_session):
|
||||
return BulkUD._factory({
|
||||
"evaluate": BulkDeleteEvaluate,
|
||||
"fetch": BulkDeleteFetch,
|
||||
False: BulkDelete
|
||||
}, synchronize_session, query)
|
||||
|
||||
def _do_exec(self):
|
||||
delete_stmt = sql.delete(self.primary_table,
|
||||
self.context.whereclause)
|
||||
|
||||
self.result = self.query.session.execute(delete_stmt,
|
||||
params=self.query._params)
|
||||
self.rowcount = self.result.rowcount
|
||||
|
||||
def _do_post(self):
|
||||
session = self.query.session
|
||||
session.dispatch.after_bulk_delete(self)
|
||||
|
||||
|
||||
class BulkUpdateEvaluate(BulkEvaluate, BulkUpdate):
|
||||
"""BulkUD which handles UPDATEs using the "evaluate"
|
||||
method of session resolution."""
|
||||
|
||||
def _additional_evaluators(self, evaluator_compiler):
|
||||
self.value_evaluators = {}
|
||||
for key, value in self.values.items():
|
||||
key = _attr_as_key(key)
|
||||
self.value_evaluators[key] = evaluator_compiler.process(
|
||||
expression._literal_as_binds(value))
|
||||
|
||||
def _do_post_synchronize(self):
|
||||
session = self.query.session
|
||||
states = set()
|
||||
evaluated_keys = list(self.value_evaluators.keys())
|
||||
for obj in self.matched_objects:
|
||||
state, dict_ = attributes.instance_state(obj),\
|
||||
attributes.instance_dict(obj)
|
||||
|
||||
# only evaluate unmodified attributes
|
||||
to_evaluate = state.unmodified.intersection(
|
||||
evaluated_keys)
|
||||
for key in to_evaluate:
|
||||
dict_[key] = self.value_evaluators[key](obj)
|
||||
|
||||
state._commit(dict_, list(to_evaluate))
|
||||
|
||||
# expire attributes with pending changes
|
||||
# (there was no autoflush, so they are overwritten)
|
||||
state._expire_attributes(dict_,
|
||||
set(evaluated_keys).
|
||||
difference(to_evaluate))
|
||||
states.add(state)
|
||||
session._register_altered(states)
|
||||
|
||||
|
||||
class BulkDeleteEvaluate(BulkEvaluate, BulkDelete):
|
||||
"""BulkUD which handles DELETEs using the "evaluate"
|
||||
method of session resolution."""
|
||||
|
||||
def _do_post_synchronize(self):
|
||||
self.query.session._remove_newly_deleted(
|
||||
[attributes.instance_state(obj)
|
||||
for obj in self.matched_objects])
|
||||
|
||||
|
||||
class BulkUpdateFetch(BulkFetch, BulkUpdate):
|
||||
"""BulkUD which handles UPDATEs using the "fetch"
|
||||
method of session resolution."""
|
||||
|
||||
def _do_post_synchronize(self):
|
||||
session = self.query.session
|
||||
target_mapper = self.query._mapper_zero()
|
||||
|
||||
states = set([
|
||||
attributes.instance_state(session.identity_map[identity_key])
|
||||
for identity_key in [
|
||||
target_mapper.identity_key_from_primary_key(
|
||||
list(primary_key))
|
||||
for primary_key in self.matched_rows
|
||||
]
|
||||
if identity_key in session.identity_map
|
||||
])
|
||||
attrib = [_attr_as_key(k) for k in self.values]
|
||||
for state in states:
|
||||
session._expire_state(state, attrib)
|
||||
session._register_altered(states)
|
||||
|
||||
|
||||
class BulkDeleteFetch(BulkFetch, BulkDelete):
|
||||
"""BulkUD which handles DELETEs using the "fetch"
|
||||
method of session resolution."""
|
||||
|
||||
def _do_post_synchronize(self):
|
||||
session = self.query.session
|
||||
target_mapper = self.query._mapper_zero()
|
||||
for primary_key in self.matched_rows:
|
||||
# TODO: inline this and call remove_newly_deleted
|
||||
# once
|
||||
identity_key = target_mapper.identity_key_from_primary_key(
|
||||
list(primary_key))
|
||||
if identity_key in session.identity_map:
|
||||
session._remove_newly_deleted(
|
||||
[attributes.instance_state(
|
||||
session.identity_map[identity_key]
|
||||
)]
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user