summaryrefslogtreecommitdiff
path: root/lib/sqlalchemy/orm
diff options
context:
space:
mode:
Diffstat (limited to 'lib/sqlalchemy/orm')
-rw-r--r--lib/sqlalchemy/orm/persistence.py12
-rw-r--r--lib/sqlalchemy/orm/query.py54
2 files changed, 42 insertions, 24 deletions
diff --git a/lib/sqlalchemy/orm/persistence.py b/lib/sqlalchemy/orm/persistence.py
index e553f399d..c3b2d7bcb 100644
--- a/lib/sqlalchemy/orm/persistence.py
+++ b/lib/sqlalchemy/orm/persistence.py
@@ -1030,6 +1030,7 @@ class BulkUD(object):
def __init__(self, query):
self.query = query.enable_eagerloads(False)
+ self.mapper = self.query._bind_mapper()
@property
def session(self):
@@ -1124,6 +1125,7 @@ class BulkFetch(BulkUD):
self.primary_table.primary_key)
self.matched_rows = session.execute(
select_stmt,
+ mapper=self.mapper,
params=query._params).fetchall()
@@ -1134,7 +1136,6 @@ class BulkUpdate(BulkUD):
super(BulkUpdate, self).__init__(query)
self.query._no_select_modifiers("update")
self.values = values
- self.mapper = self.query._mapper_zero_or_none()
@classmethod
def factory(cls, query, synchronize_session, values):
@@ -1180,7 +1181,8 @@ class BulkUpdate(BulkUD):
self.context.whereclause, values)
self.result = self.query.session.execute(
- update_stmt, params=self.query._params)
+ update_stmt, params=self.query._params,
+ mapper=self.mapper)
self.rowcount = self.result.rowcount
def _do_post(self):
@@ -1207,8 +1209,10 @@ class BulkDelete(BulkUD):
delete_stmt = sql.delete(self.primary_table,
self.context.whereclause)
- self.result = self.query.session.execute(delete_stmt,
- params=self.query._params)
+ self.result = self.query.session.execute(
+ delete_stmt,
+ params=self.query._params,
+ mapper=self.mapper)
self.rowcount = self.result.rowcount
def _do_post(self):
diff --git a/lib/sqlalchemy/orm/query.py b/lib/sqlalchemy/orm/query.py
index 7302574e6..cd8b0efbe 100644
--- a/lib/sqlalchemy/orm/query.py
+++ b/lib/sqlalchemy/orm/query.py
@@ -146,7 +146,7 @@ class Query(object):
ext_info,
aliased_adapter
)
- ent.setup_entity(*d[entity])
+ ent.setup_entity(ent, *d[entity])
def _mapper_loads_polymorphically_with(self, mapper, adapter):
for m2 in mapper._with_polymorphic_mappers or [mapper]:
@@ -160,7 +160,6 @@ class Query(object):
for from_obj in obj:
info = inspect(from_obj)
-
if hasattr(info, 'mapper') and \
(info.is_mapper or info.is_aliased_class):
self._select_from_entity = from_obj
@@ -286,8 +285,9 @@ class Query(object):
return self._entities[0]
def _mapper_zero(self):
- return self._select_from_entity or \
- self._entity_zero().entity_zero
+ return self._select_from_entity \
+ if self._select_from_entity is not None \
+ else self._entity_zero().entity_zero
@property
def _mapper_entities(self):
@@ -301,11 +301,14 @@ class Query(object):
self._mapper_zero()
)
- def _mapper_zero_or_none(self):
- if self._primary_entity:
- return self._primary_entity.mapper
- else:
- return None
+ def _bind_mapper(self):
+ ezero = self._mapper_zero()
+ if ezero is not None:
+ insp = inspect(ezero)
+ if hasattr(insp, 'mapper'):
+ return insp.mapper
+
+ return None
def _only_mapper_zero(self, rationale=None):
if len(self._entities) > 1:
@@ -988,6 +991,7 @@ class Query(object):
statement.correlate(None)
q = self._from_selectable(fromclause)
q._enable_single_crit = False
+ q._select_from_entity = self._mapper_zero()
if entities:
q._set_entities(entities)
return q
@@ -2526,7 +2530,7 @@ class Query(object):
def _execute_and_instances(self, querycontext):
conn = self._connection_from_session(
- mapper=self._mapper_zero_or_none(),
+ mapper=self._bind_mapper(),
clause=querycontext.statement,
close_with_result=True)
@@ -3160,7 +3164,7 @@ class _MapperEntity(_QueryEntity):
supports_single_entity = True
- def setup_entity(self, ext_info, aliased_adapter):
+ def setup_entity(self, original_entity, ext_info, aliased_adapter):
self.mapper = ext_info.mapper
self.aliased_adapter = aliased_adapter
self.selectable = ext_info.selectable
@@ -3507,9 +3511,9 @@ class _BundleEntity(_QueryEntity):
for ent in self._entities:
ent.adapt_to_selectable(c, sel)
- def setup_entity(self, ext_info, aliased_adapter):
+ def setup_entity(self, original_entity, ext_info, aliased_adapter):
for ent in self._entities:
- ent.setup_entity(ext_info, aliased_adapter)
+ ent.setup_entity(original_entity, ext_info, aliased_adapter)
def setup_context(self, query, context):
for ent in self._entities:
@@ -3592,15 +3596,23 @@ class _ColumnEntity(_QueryEntity):
# leaking out their entities into the main select construct
self.actual_froms = actual_froms = set(column._from_objects)
- self.entities = util.OrderedSet(
- elem._annotations['parententity']
- for elem in visitors.iterate(column, {})
+ all_elements = [
+ elem for elem in visitors.iterate(column, {})
if 'parententity' in elem._annotations
- and actual_froms.intersection(elem._from_objects)
+ ]
+
+ self.entities = util.unique_list([
+ elem._annotations['parententity']
+ for elem in all_elements
+ ])
+ self._from_entities = set(
+ elem._annotations['parententity']
+ for elem in all_elements
+ if actual_froms.intersection(elem._from_objects)
)
if self.entities:
- self.entity_zero = list(self.entities)[0]
+ self.entity_zero = self.entities[0]
elif self.namespace is not None:
self.entity_zero = self.namespace
else:
@@ -3623,10 +3635,12 @@ class _ColumnEntity(_QueryEntity):
c.entity_zero = self.entity_zero
c.entities = self.entities
- def setup_entity(self, ext_info, aliased_adapter):
+ def setup_entity(self, original_entity, ext_info, aliased_adapter):
if 'selectable' not in self.__dict__:
self.selectable = ext_info.selectable
- self.froms.add(ext_info.selectable)
+
+ if original_entity in self._from_entities:
+ self.froms.add(ext_info.selectable)
def corresponds_to(self, entity):
# TODO: just returning False here,