diff options
Diffstat (limited to 'lib')
| -rw-r--r-- | lib/sqlalchemy/ext/sqlsoup.py | 2 | ||||
| -rw-r--r-- | lib/sqlalchemy/orm/__init__.py | 16 | ||||
| -rw-r--r-- | lib/sqlalchemy/orm/attributes.py | 98 | ||||
| -rw-r--r-- | lib/sqlalchemy/orm/dependency.py | 112 | ||||
| -rw-r--r-- | lib/sqlalchemy/orm/dynamic.py | 16 | ||||
| -rw-r--r-- | lib/sqlalchemy/orm/interfaces.py | 6 | ||||
| -rw-r--r-- | lib/sqlalchemy/orm/mapper.py | 85 | ||||
| -rw-r--r-- | lib/sqlalchemy/orm/properties.py | 290 | ||||
| -rw-r--r-- | lib/sqlalchemy/orm/query.py | 4 | ||||
| -rw-r--r-- | lib/sqlalchemy/orm/session.py | 2 | ||||
| -rw-r--r-- | lib/sqlalchemy/orm/strategies.py | 180 | ||||
| -rw-r--r-- | lib/sqlalchemy/orm/unitofwork.py | 17 | ||||
| -rw-r--r-- | lib/sqlalchemy/orm/util.py | 4 | ||||
| -rw-r--r-- | lib/sqlalchemy/schema.py | 10 |
14 files changed, 496 insertions, 346 deletions
diff --git a/lib/sqlalchemy/ext/sqlsoup.py b/lib/sqlalchemy/ext/sqlsoup.py index f306b1559..fbbf8d6fd 100644 --- a/lib/sqlalchemy/ext/sqlsoup.py +++ b/lib/sqlalchemy/ext/sqlsoup.py @@ -419,7 +419,7 @@ class TableClassType(SelectableClassType): cls._table.update(whereclause, values).execute(**kwargs) def relate(cls, propname, *args, **kwargs): - class_mapper(cls)._compile_property(propname, relation(*args, **kwargs)) + class_mapper(cls)._configure_property(propname, relation(*args, **kwargs)) def _is_outer_join(selectable): if not isinstance(selectable, sql.Join): diff --git a/lib/sqlalchemy/orm/__init__.py b/lib/sqlalchemy/orm/__init__.py index 3c539b8f4..769c24886 100644 --- a/lib/sqlalchemy/orm/__init__.py +++ b/lib/sqlalchemy/orm/__init__.py @@ -201,11 +201,19 @@ def relation(argument, secondary=None, **kwargs): keyword argument. :param backref: - indicates the name of a property to be placed on the related + indicates the string name of a property to be placed on the related mapper's class that will handle this relationship in the other - direction, including synchronizing the object attributes on both - sides of the relation. Can also point to a :func:`backref` for - more configurability. + direction. The other property will be created automatically + when the mappers are configured. Can also be passed as a + :func:`backref` object to control the configuration of the + new relation. + + :param back_populates: + Takes a string name and has the same meaning as ``backref``, + except the complementing property is **not** created automatically, + and instead must be configured explicitly on the other mapper. The + complementing property should also indicate ``back_populates`` + to this relation to ensure proper functioning. :param cascade: a comma-separated list of cascade rules which determines how diff --git a/lib/sqlalchemy/orm/attributes.py b/lib/sqlalchemy/orm/attributes.py index 5c242aa7e..1606f8267 100644 --- a/lib/sqlalchemy/orm/attributes.py +++ b/lib/sqlalchemy/orm/attributes.py @@ -93,7 +93,7 @@ ClassManager instrumentation is used. class QueryableAttribute(interfaces.PropComparator): - def __init__(self, impl, comparator=None, parententity=None): + def __init__(self, key, impl=None, comparator=None, parententity=None): """Construct an InstrumentedAttribute. comparator @@ -104,12 +104,6 @@ class QueryableAttribute(interfaces.PropComparator): self.comparator = comparator self.parententity = parententity - if parententity: - mapper, selectable, is_aliased_class = _entity_info(parententity, compile=False) - self.property = mapper._get_property(self.impl.key) - else: - self.property = None - def get_history(self, instance, **kwargs): return self.impl.get_history(instance_state(instance), **kwargs) @@ -145,6 +139,11 @@ class QueryableAttribute(interfaces.PropComparator): def __str__(self): return repr(self.parententity) + "." + self.property.key + @property + def property(self): + return self.comparator.property + + class InstrumentedAttribute(QueryableAttribute): """Public-facing descriptor, placed in the mapped class dictionary.""" @@ -234,13 +233,13 @@ class AttributeImpl(object): """internal implementation for instrumented attributes.""" def __init__(self, class_, key, - callable_, class_manager, trackparent=False, extension=None, - compare_function=None, active_history=False, **kwargs): + callable_, trackparent=False, extension=None, + compare_function=None, active_history=False, parent_token=None, **kwargs): """Construct an AttributeImpl. \class_ - the class to be instrumented. - + associated class + key string name of the attribute @@ -267,13 +266,18 @@ class AttributeImpl(object): even if it means executing a lazy callable upon attribute change. This flag is set to True if any extensions are present. + parent_token + Usually references the MapperProperty, used as a key for + the hasparent() function to identify an "owning" attribute. + Allows multiple AttributeImpls to all match a single + owner attribute. + """ - self.class_ = class_ self.key = key self.callable_ = callable_ - self.class_manager = class_manager self.trackparent = trackparent + self.parent_token = parent_token or self if compare_function is None: self.is_equal = operator.eq else: @@ -296,7 +300,7 @@ class AttributeImpl(object): will also not have a `hasparent` flag. """ - return state.parents.get(id(self), optimistic) + return state.parents.get(id(self.parent_token), optimistic) def sethasparent(self, state, value): """Set a boolean flag on the given item corresponding to @@ -304,7 +308,7 @@ class AttributeImpl(object): attribute represented by this ``InstrumentedAttribute``. """ - state.parents[id(self)] = value + state.parents[id(self.parent_token)] = value def set_callable(self, state, callable_): """Set a callable function for this attribute on the given object. @@ -472,7 +476,7 @@ class MutableScalarAttributeImpl(ScalarAttributeImpl): class_manager, copy_function=None, compare_function=None, **kwargs): super(ScalarAttributeImpl, self).__init__(class_, key, callable_, - class_manager, compare_function=compare_function, **kwargs) + compare_function=compare_function, **kwargs) class_manager.mutable_attributes.add(key) if copy_function is None: raise sa_exc.ArgumentError("MutableScalarAttributeImpl requires a copy function") @@ -512,11 +516,11 @@ class ScalarObjectAttributeImpl(ScalarAttributeImpl): accepts_scalar_loader = False uses_objects = True - def __init__(self, class_, key, callable_, class_manager, + def __init__(self, class_, key, callable_, trackparent=False, extension=None, copy_function=None, compare_function=None, **kwargs): super(ScalarObjectAttributeImpl, self).__init__(class_, key, - callable_, class_manager, trackparent=trackparent, extension=extension, + callable_, trackparent=trackparent, extension=extension, compare_function=compare_function, **kwargs) if compare_function is None: self.is_equal = identity_equal @@ -589,11 +593,10 @@ class CollectionAttributeImpl(AttributeImpl): accepts_scalar_loader = False uses_objects = True - def __init__(self, class_, key, callable_, class_manager, + def __init__(self, class_, key, callable_, typecallable=None, trackparent=False, extension=None, copy_function=None, compare_function=None, **kwargs): - super(CollectionAttributeImpl, self).__init__(class_, - key, callable_, class_manager, trackparent=trackparent, + super(CollectionAttributeImpl, self).__init__(class_, key, callable_, trackparent=trackparent, extension=extension, compare_function=compare_function, **kwargs) if copy_function is None: @@ -1540,36 +1543,37 @@ def unregister_class(class_): manager.instantiable = False manager.unregister() -def register_attribute(class_, key, uselist, useobject, - callable_=None, proxy_property=None, - mutable_scalars=False, impl_class=None, **kwargs): - manager = manager_of_class(class_) - if manager.is_instrumented(key): - return +def register_attribute(class_, key, **kw): + proxy_property = kw.pop('proxy_property', None) + + comparator = kw.pop('comparator', None) + parententity = kw.pop('parententity', None) + register_descriptor(class_, key, proxy_property, comparator, parententity) + if not proxy_property: + register_attribute_impl(class_, key, **kw) + +def register_attribute_impl(class_, key, **kw): + + manager = manager_of_class(class_) + uselist = kw.get('uselist', False) if uselist: - factory = kwargs.pop('typecallable', None) + factory = kw.pop('typecallable', None) typecallable = manager.instrument_collection_class( key, factory or list) else: - typecallable = kwargs.pop('typecallable', None) + typecallable = kw.pop('typecallable', None) + + manager[key].impl = _create_prop(class_, key, manager, typecallable=typecallable, **kw) - comparator = kwargs.pop('comparator', None) - parententity = kwargs.pop('parententity', None) +def register_descriptor(class_, key, proxy_property=None, comparator=None, parententity=None, property_=None): + manager = manager_of_class(class_) if proxy_property: proxy_type = proxied_attribute_factory(proxy_property) descriptor = proxy_type(key, proxy_property, comparator, parententity) else: - descriptor = InstrumentedAttribute( - _create_prop(class_, key, uselist, callable_, - class_manager=manager, - useobject=useobject, - typecallable=typecallable, - mutable_scalars=mutable_scalars, - impl_class=impl_class, - **kwargs), - comparator=comparator, parententity=parententity) + descriptor = InstrumentedAttribute(key, comparator=comparator, parententity=parententity) manager.instrument_attribute(key, descriptor) @@ -1741,22 +1745,24 @@ def collect_management_factories_for(cls): factories.discard(None) return factories -def _create_prop(class_, key, uselist, callable_, class_manager, typecallable, useobject, mutable_scalars, impl_class, **kwargs): +def _create_prop(class_, key, class_manager, + uselist=False, callable_=None, typecallable=None, + useobject=False, mutable_scalars=False, + impl_class=None, **kwargs): if impl_class: - return impl_class(class_, key, typecallable, class_manager=class_manager, **kwargs) + return impl_class(class_, key, typecallable, **kwargs) elif uselist: return CollectionAttributeImpl(class_, key, callable_, typecallable=typecallable, - class_manager=class_manager, **kwargs) + **kwargs) elif useobject: return ScalarObjectAttributeImpl(class_, key, callable_, - class_manager=class_manager, **kwargs) + **kwargs) elif mutable_scalars: return MutableScalarAttributeImpl(class_, key, callable_, class_manager=class_manager, **kwargs) else: - return ScalarAttributeImpl(class_, key, callable_, - class_manager=class_manager, **kwargs) + return ScalarAttributeImpl(class_, key, callable_, **kwargs) def _generate_init(class_, class_manager): """Build an __init__ decorator that triggers ClassManager events.""" diff --git a/lib/sqlalchemy/orm/dependency.py b/lib/sqlalchemy/orm/dependency.py index fb24f6a68..516295709 100644 --- a/lib/sqlalchemy/orm/dependency.py +++ b/lib/sqlalchemy/orm/dependency.py @@ -36,7 +36,6 @@ class DependencyProcessor(object): self.parent = prop.parent self.secondary = prop.secondary self.direction = prop.direction - self.is_backref = prop._is_backref self.post_update = prop.post_update self.passive_deletes = prop.passive_deletes self.passive_updates = prop.passive_updates @@ -44,19 +43,21 @@ class DependencyProcessor(object): self.key = prop.key self.dependency_marker = MapperStub(self.parent, self.mapper, self.key) if not self.prop.synchronize_pairs: - raise sa_exc.ArgumentError("Can't build a DependencyProcessor for relation %s. No target attributes to populate between parent and child are present" % self.prop) + raise sa_exc.ArgumentError("Can't build a DependencyProcessor for relation %s. " + "No target attributes to populate between parent and child are present" % self.prop) def _get_instrumented_attribute(self): """Return the ``InstrumentedAttribute`` handled by this ``DependencyProecssor``. + """ - return self.parent.class_manager.get_impl(self.key) def hasparent(self, state): """return True if the given object instance has a parent, - according to the ``InstrumentedAttribute`` handled by this ``DependencyProcessor``.""" - + according to the ``InstrumentedAttribute`` handled by this ``DependencyProcessor``. + + """ # TODO: use correct API for this return self._get_instrumented_attribute().hasparent(state) @@ -78,8 +79,8 @@ class DependencyProcessor(object): """Given an object pair assuming `obj2` is a child of `obj1`, return a tuple with the dependent object second, or None if there is no dependency. - """ + """ if state1 is state2: return None elif self.direction == ONETOMANY: @@ -94,8 +95,8 @@ class DependencyProcessor(object): It is called within the context of the various mappers and sometimes individual objects sorted according to their insert/update/delete order (topological sort). - """ + """ raise NotImplementedError() def preprocess_dependencies(self, task, deplist, uowcommit, delete = False): @@ -103,26 +104,51 @@ class DependencyProcessor(object): through related objects and ensure every instance which will require save/update/delete is properly added to the UOWTransaction. - """ + """ raise NotImplementedError() def _verify_canload(self, state): if state is not None and not self.mapper._canload(state, allow_subtypes=not self.enable_typechecks): if self.mapper._canload(state, allow_subtypes=True): - raise exc.FlushError("Attempting to flush an item of type %s on collection '%s', which is not the expected type %s. Configure mapper '%s' to load this subtype polymorphically, or set enable_typechecks=False to allow subtypes. Mismatched typeloading may cause bi-directional relationships (backrefs) to not function properly." % (state.class_, self.prop, self.mapper.class_, self.mapper)) + raise exc.FlushError("Attempting to flush an item of type %s on collection '%s', " + "which is not the expected type %s. Configure mapper '%s' to load this " + "subtype polymorphically, or set enable_typechecks=False to allow subtypes. " + "Mismatched typeloading may cause bi-directional relationships (backrefs) " + "to not function properly." % (state.class_, self.prop, self.mapper.class_, self.mapper)) else: - raise exc.FlushError("Attempting to flush an item of type %s on collection '%s', whose mapper does not inherit from that of %s." % (state.class_, self.prop, self.mapper.class_)) + raise exc.FlushError("Attempting to flush an item of type %s on collection '%s', " + "whose mapper does not inherit from that of %s." % (state.class_, self.prop, self.mapper.class_)) def _synchronize(self, state, child, associationrow, clearkeys, uowcommit): """Called during a flush to synchronize primary key identifier values between a parent/child object, as well as to an associationrow in the case of many-to-many. + """ - raise NotImplementedError() - + def _check_reverse_action(self, uowcommit, parent, child, action): + """Determine if an action has been performed by the 'reverse' property of this property. + + this is used to ensure that only one side of a bidirectional relation + issues a certain operation for a parent/child pair. + + """ + for r in self.prop._reverse_property: + if (r._dependency_processor, action, parent, child) in uowcommit.attributes: + return True + return False + + def _performed_action(self, uowcommit, parent, child, action): + """Establish that an action has been performed for a certain parent/child pair. + + Used only for actions that are sensitive to bidirectional double-action, + i.e. manytomany, post_update. + + """ + uowcommit.attributes[(self, action, parent, child)] = True + def _conditional_post_update(self, state, uowcommit, related): """Execute a post_update call. @@ -135,33 +161,32 @@ class DependencyProcessor(object): particular relationship, and given a target object and list of one or more related objects, and execute the ``UPDATE`` if the given related object list contains ``INSERT``s or ``DELETE``s. + """ - if state is not None and self.post_update: for x in related: - if x is not None: + if x is not None and not self._check_reverse_action(uowcommit, x, state, "postupdate"): uowcommit.register_object(state, postupdate=True, post_update_cols=[r for l, r in self.prop.synchronize_pairs]) + self._performed_action(uowcommit, x, state, "postupdate") break def _pks_changed(self, uowcommit, state): raise NotImplementedError() def __repr__(self): - return "%s(%s)" % (self.__class__.__name__, str(self.prop)) + return "%s(%s)" % (self.__class__.__name__, self.prop) class OneToManyDP(DependencyProcessor): def register_dependencies(self, uowcommit): if self.post_update: - if not self.is_backref: - uowcommit.register_dependency(self.mapper, self.dependency_marker) - uowcommit.register_dependency(self.parent, self.dependency_marker) - uowcommit.register_processor(self.dependency_marker, self, self.parent) + uowcommit.register_dependency(self.mapper, self.dependency_marker) + uowcommit.register_dependency(self.parent, self.dependency_marker) + uowcommit.register_processor(self.dependency_marker, self, self.parent) else: uowcommit.register_dependency(self.parent, self.mapper) uowcommit.register_processor(self.parent, self, self.parent) def process_dependencies(self, task, deplist, uowcommit, delete = False): - #print self.mapper.mapped_table.name + " " + self.key + " " + repr(len(deplist)) + " process_dep isdelete " + repr(delete) + " direction " + repr(self.direction) if delete: # head object is being deleted, and we manage its list of child objects # the child objects have to have their foreign key to the parent set to NULL @@ -198,8 +223,6 @@ class OneToManyDP(DependencyProcessor): self._synchronize(state, child, None, False, uowcommit) def preprocess_dependencies(self, task, deplist, uowcommit, delete = False): - #print self.mapper.mapped_table.name + " " + self.key + " " + repr(len(deplist)) + " preprocess_dep isdelete " + repr(delete) + " direction " + repr(self.direction) - if delete: # head object is being deleted, and we manage its list of child objects # the child objects have to have their foreign key to the parent set to NULL @@ -304,17 +327,15 @@ class ManyToOneDP(DependencyProcessor): def register_dependencies(self, uowcommit): if self.post_update: - if not self.is_backref: - uowcommit.register_dependency(self.mapper, self.dependency_marker) - uowcommit.register_dependency(self.parent, self.dependency_marker) - uowcommit.register_processor(self.dependency_marker, self, self.parent) + uowcommit.register_dependency(self.mapper, self.dependency_marker) + uowcommit.register_dependency(self.parent, self.dependency_marker) + uowcommit.register_processor(self.dependency_marker, self, self.parent) else: uowcommit.register_dependency(self.mapper, self.parent) uowcommit.register_processor(self.mapper, self, self.parent) def process_dependencies(self, task, deplist, uowcommit, delete=False): - #print self.mapper.mapped_table.name + " " + self.key + " " + repr(len(deplist)) + " process_dep isdelete " + repr(delete) + " direction " + repr(self.direction) if delete: if self.post_update and not self.cascade.delete_orphan and not self.passive_deletes == 'all': # post_update means we have to update our row to not reference the child object @@ -333,7 +354,6 @@ class ManyToOneDP(DependencyProcessor): self._conditional_post_update(state, uowcommit, history.sum()) def preprocess_dependencies(self, task, deplist, uowcommit, delete=False): - #print self.mapper.mapped_table.name + " " + self.key + " " + repr(len(deplist)) + " PRE process_dep isdelete " + repr(delete) + " direction " + repr(self.direction) if self.post_update: return if delete: @@ -390,45 +410,39 @@ class ManyToManyDP(DependencyProcessor): uowcommit.register_processor(self.dependency_marker, self, self.parent) def process_dependencies(self, task, deplist, uowcommit, delete = False): - #print self.mapper.mapped_table.name + " " + self.key + " " + repr(len(deplist)) + " process_dep isdelete " + repr(delete) + " direction " + repr(self.direction) connection = uowcommit.transaction.connection(self.mapper) secondary_delete = [] secondary_insert = [] secondary_update = [] - if self.prop._reverse_property: - reverse_dep = getattr(self.prop._reverse_property, '_dependency_processor', None) - else: - reverse_dep = None - if delete: for state in deplist: history = uowcommit.get_attribute_history(state, self.key, passive=self.passive_deletes) if history: for child in history.non_added(): - if child is None or (reverse_dep and (reverse_dep, "manytomany", child, state) in uowcommit.attributes): + if child is None or self._check_reverse_action(uowcommit, child, state, "manytomany"): continue associationrow = {} self._synchronize(state, child, associationrow, False, uowcommit) secondary_delete.append(associationrow) - uowcommit.attributes[(self, "manytomany", state, child)] = True + self._performed_action(uowcommit, state, child, "manytomany") else: for state in deplist: history = uowcommit.get_attribute_history(state, self.key) if history: for child in history.added: - if child is None or (reverse_dep and (reverse_dep, "manytomany", child, state) in uowcommit.attributes): + if child is None or self._check_reverse_action(uowcommit, child, state, "manytomany"): continue associationrow = {} self._synchronize(state, child, associationrow, False, uowcommit) - uowcommit.attributes[(self, "manytomany", state, child)] = True + self._performed_action(uowcommit, state, child, "manytomany") secondary_insert.append(associationrow) for child in history.deleted: - if child is None or (reverse_dep and (reverse_dep, "manytomany", child, state) in uowcommit.attributes): + if child is None or self._check_reverse_action(uowcommit, child, state, "manytomany"): continue associationrow = {} self._synchronize(state, child, associationrow, False, uowcommit) - uowcommit.attributes[(self, "manytomany", state, child)] = True + self._performed_action(uowcommit, state, child, "manytomany") secondary_delete.append(associationrow) if not self.passive_updates and self._pks_changed(uowcommit, state): @@ -444,24 +458,30 @@ class ManyToManyDP(DependencyProcessor): secondary_update.append(associationrow) if secondary_delete: - # TODO: precompile the delete/insert queries? - statement = self.secondary.delete(sql.and_(*[c == sql.bindparam(c.key, type_=c.type) for c in self.secondary.c if c.key in associationrow])) + statement = self.secondary.delete(sql.and_(*[ + c == sql.bindparam(c.key, type_=c.type) for c in self.secondary.c if c.key in associationrow + ])) result = connection.execute(statement, secondary_delete) if result.supports_sane_multi_rowcount() and result.rowcount != len(secondary_delete): - raise exc.ConcurrentModificationError("Deleted rowcount %d does not match number of secondary table rows deleted from table '%s': %d" % (result.rowcount, self.secondary.description, len(secondary_delete))) + raise exc.ConcurrentModificationError("Deleted rowcount %d does not match number of " + "secondary table rows deleted from table '%s': %d" % + (result.rowcount, self.secondary.description, len(secondary_delete))) if secondary_update: - statement = self.secondary.update(sql.and_(*[c == sql.bindparam("old_" + c.key, type_=c.type) for c in self.secondary.c if c.key in associationrow])) + statement = self.secondary.update(sql.and_(*[ + c == sql.bindparam("old_" + c.key, type_=c.type) for c in self.secondary.c if c.key in associationrow + ])) result = connection.execute(statement, secondary_update) if result.supports_sane_multi_rowcount() and result.rowcount != len(secondary_update): - raise exc.ConcurrentModificationError("Updated rowcount %d does not match number of secondary table rows updated from table '%s': %d" % (result.rowcount, self.secondary.description, len(secondary_update))) + raise exc.ConcurrentModificationError("Updated rowcount %d does not match number of " + "secondary table rows updated from table '%s': %d" % + (result.rowcount, self.secondary.description, len(secondary_update))) if secondary_insert: statement = self.secondary.insert() connection.execute(statement, secondary_insert) def preprocess_dependencies(self, task, deplist, uowcommit, delete = False): - #print self.mapper.mapped_table.name + " " + self.key + " " + repr(len(deplist)) + " preprocess_dep isdelete " + repr(delete) + " direction " + repr(self.direction) if not delete: for state in deplist: history = uowcommit.get_attribute_history(state, self.key, passive=True) diff --git a/lib/sqlalchemy/orm/dynamic.py b/lib/sqlalchemy/orm/dynamic.py index 1bc0994c1..a46734dde 100644 --- a/lib/sqlalchemy/orm/dynamic.py +++ b/lib/sqlalchemy/orm/dynamic.py @@ -24,7 +24,14 @@ from sqlalchemy.orm.util import _state_has_identity, has_identity class DynaLoader(strategies.AbstractRelationLoader): def init_class_attribute(self): self.is_class_level = True - self._register_attribute(self.parent.class_, impl_class=DynamicAttributeImpl, target_mapper=self.parent_property.mapper, order_by=self.parent_property.order_by, query_class=self.parent_property.query_class) + + strategies._register_attribute(self, + useobject=True, + impl_class=DynamicAttributeImpl, + target_mapper=self.parent_property.mapper, + order_by=self.parent_property.order_by, + query_class=self.parent_property.query_class + ) def create_row_processor(self, selectcontext, path, mapper, row, adapter): return (None, None) @@ -35,10 +42,9 @@ class DynamicAttributeImpl(attributes.AttributeImpl): uses_objects = True accepts_scalar_loader = False - def __init__(self, class_, key, typecallable, class_manager, - target_mapper, order_by, query_class=None, **kwargs): - super(DynamicAttributeImpl, self).__init__( - class_, key, typecallable, class_manager, **kwargs) + def __init__(self, class_, key, typecallable, + target_mapper, order_by, query_class=None, **kwargs): + super(DynamicAttributeImpl, self).__init__(class_, key, typecallable, **kwargs) self.target_mapper = target_mapper self.order_by = order_by if not query_class: diff --git a/lib/sqlalchemy/orm/interfaces.py b/lib/sqlalchemy/orm/interfaces.py index b210e577f..fb77f56c2 100644 --- a/lib/sqlalchemy/orm/interfaces.py +++ b/lib/sqlalchemy/orm/interfaces.py @@ -392,13 +392,15 @@ class MapperProperty(object): def set_parent(self, parent): self.parent = parent - def init(self, key, parent): + def instrument_class(self, mapper): + raise NotImplementedError() + + def init(self): """Called after all mappers are compiled to assemble relationships between mappers, establish instrumented class attributes. """ - self.key = key self._compiled = True self.do_init() diff --git a/lib/sqlalchemy/orm/mapper.py b/lib/sqlalchemy/orm/mapper.py index 96043972a..17a12e70f 100644 --- a/lib/sqlalchemy/orm/mapper.py +++ b/lib/sqlalchemy/orm/mapper.py @@ -58,6 +58,8 @@ _COMPILE_MUTEX = util.threading.RLock() ColumnProperty = None SynonymProperty = None ComparableProperty = None +RelationProperty = None +ConcreteInheritedProperty = None _expire_state = None _state_session = None @@ -113,7 +115,7 @@ class Mapper(object): self.order_by = util.to_list(order_by) else: self.order_by = order_by - + self.always_refresh = always_refresh self.version_id_col = version_id_col self.concrete = concrete @@ -476,13 +478,13 @@ class Mapper(object): # load custom properties if self._init_properties: for key, prop in self._init_properties.iteritems(): - self._compile_property(key, prop, False) + self._configure_property(key, prop, False) # pull properties from the inherited mapper if any. if self.inherits: for key, prop in self.inherits._props.iteritems(): if key not in self._props and not self._should_exclude(key, local=False): - self._adapt_inherited_property(key, prop) + self._adapt_inherited_property(key, prop, False) # create properties for each column in the mapped table, # for those columns which don't already map to a property @@ -501,53 +503,29 @@ class Mapper(object): if column in mapper._columntoproperty: column_key = mapper._columntoproperty[column].key - self._compile_property(column_key, column, init=False, setparent=True) + self._configure_property(column_key, column, init=False, setparent=True) # do a special check for the "discriminiator" column, as it may only be present # in the 'with_polymorphic' selectable but we need it for the base mapper if self.polymorphic_on and self.polymorphic_on not in self._columntoproperty: - col = self.mapped_table.corresponding_column(self.polymorphic_on) or self.polymorphic_on + col = self.mapped_table.corresponding_column(self.polymorphic_on) + if not col: + dont_instrument = True + col = self.polymorphic_on + else: + dont_instrument = False if self._should_exclude(col.key, local=False): raise sa_exc.InvalidRequestError("Cannot exclude or override the discriminator column %r" % col.key) - self._compile_property(col.key, ColumnProperty(col), init=False, setparent=True) + self._configure_property(col.key, ColumnProperty(col, _no_instrument=dont_instrument), init=False, setparent=True) - def _adapt_inherited_property(self, key, prop): + def _adapt_inherited_property(self, key, prop, init): if not self.concrete: - self._compile_property(key, prop, init=False, setparent=False) - # TODO: concrete properties dont adapt at all right now....will require copies of relations() etc. - - class _CompileOnAttr(PropComparator): - """A placeholder descriptor which triggers compilation on access.""" - - def __init__(self, class_, key): - self.class_ = class_ - self.key = key - self.existing_prop = getattr(class_, key, None) - - def __getattribute__(self, key): - cls = object.__getattribute__(self, 'class_') - clskey = object.__getattribute__(self, 'key') - - # ugly hack - if key.startswith('__') and key != '__clause_element__': - return object.__getattribute__(self, key) - - class_mapper(cls) - - if cls.__dict__.get(clskey) is self: - # if this warning occurs, it usually means mapper - # compilation has failed, but operations upon the mapped - # classes have proceeded. - util.warn( - ("Attribute '%s' on class '%s' was not replaced during " - "mapper compilation operation") % (clskey, cls.__name__)) - # clean us up explicitly - delattr(cls, clskey) - - return getattr(getattr(cls, clskey), key) - - def _compile_property(self, key, prop, init=True, setparent=True): - self._log("_compile_property(%s, %s)" % (key, prop.__class__.__name__)) + self._configure_property(key, prop, init=False, setparent=False) + elif key not in self._props: + self._configure_property(key, ConcreteInheritedProperty(), init=init, setparent=True) + + def _configure_property(self, key, prop, init=True, setparent=True): + self._log("_configure_property(%s, %s)" % (key, prop.__class__.__name__)) if not isinstance(prop, MapperProperty): # we were passed a Column or a list of Columns; generate a ColumnProperty @@ -568,7 +546,7 @@ class Mapper(object): prop = prop.copy() prop.columns.append(column) self._log("appending to existing ColumnProperty %s" % (key)) - elif prop is None: + elif prop is None or isinstance(prop, ConcreteInheritedProperty): mapped_column = [] for c in columns: mc = self.mapped_table.corresponding_column(c) @@ -619,8 +597,6 @@ class Mapper(object): elif isinstance(prop, (ComparableProperty, SynonymProperty)) and setparent: if prop.descriptor is None: desc = getattr(self.class_, key, None) - if isinstance(desc, Mapper._CompileOnAttr): - desc = object.__getattribute__(desc, 'existing_prop') if self._is_userland_descriptor(desc): prop.descriptor = desc if getattr(prop, 'map_column', False): @@ -628,7 +604,7 @@ class Mapper(object): raise sa_exc.ArgumentError( "Can't compile synonym '%s': no column on table '%s' named '%s'" % (prop.name, self.mapped_table.description, key)) - self._compile_property(prop.name, ColumnProperty(self.mapped_table.c[key]), init=init, setparent=setparent) + self._configure_property(prop.name, ColumnProperty(self.mapped_table.c[key]), init=init, setparent=setparent) self._props[key] = prop prop.key = key @@ -636,15 +612,15 @@ class Mapper(object): if setparent: prop.set_parent(self) - if not self.non_primary: - self.class_manager.install_descriptor( - key, Mapper._CompileOnAttr(self.class_, key)) + if not self.non_primary: + prop.instrument_class(self) + + for mapper in self._inheriting_mappers: + mapper._adapt_inherited_property(key, prop, init) if init: - prop.init(key, self) + prop.init() - for mapper in self._inheriting_mappers: - mapper._adapt_inherited_property(key, prop) def compile(self): """Compile this mapper and all other non-compiled mappers. @@ -710,7 +686,7 @@ class Mapper(object): for key, prop in l: if not getattr(prop, '_compiled', False): self._log("initialize prop " + key) - prop.init(key, self) + prop.init() self._log("_post_configure_properties() complete") self.compiled = True @@ -732,7 +708,7 @@ class Mapper(object): """ self._init_properties[key] = prop - self._compile_property(key, prop, init=self.compiled) + self._configure_property(key, prop, init=self.compiled) # class formatting / logging. @@ -828,6 +804,7 @@ class Mapper(object): construct an outerjoin amongst those mapper's mapped tables. """ + from_obj = self.mapped_table for m in mappers: if m is self: diff --git a/lib/sqlalchemy/orm/properties.py b/lib/sqlalchemy/orm/properties.py index 343b73f42..96c4565c1 100644 --- a/lib/sqlalchemy/orm/properties.py +++ b/lib/sqlalchemy/orm/properties.py @@ -41,15 +41,30 @@ class ColumnProperty(StrategizedProperty): self.columns = [expression._labeled(c) for c in columns] self.group = kwargs.pop('group', None) self.deferred = kwargs.pop('deferred', False) + self.no_instrument = kwargs.pop('_no_instrument', False) self.comparator_factory = kwargs.pop('comparator_factory', self.__class__.Comparator) self.descriptor = kwargs.pop('descriptor', None) self.extension = kwargs.pop('extension', None) util.set_creation_order(self) - if self.deferred: + if self.no_instrument: + self.strategy_class = strategies.UninstrumentedColumnLoader + elif self.deferred: self.strategy_class = strategies.DeferredColumnLoader else: self.strategy_class = strategies.ColumnLoader - + + def instrument_class(self, mapper): + if self.no_instrument: + return + + attributes.register_descriptor( + mapper.class_, + self.key, + comparator=self.comparator_factory(self, mapper), + parententity=mapper, + property_=self + ) + def do_init(self): super(ColumnProperty, self).do_init() if len(self.columns) > 1 and self.parent.primary_key.issuperset(self.columns): @@ -57,8 +72,8 @@ class ColumnProperty(StrategizedProperty): ("On mapper %s, primary key column '%s' is being combined " "with distinct primary key column '%s' in attribute '%s'. " "Use explicit properties to give each column its own mapped " - "attribute name.") % (str(self.parent), str(self.columns[1]), - str(self.columns[0]), self.key)) + "attribute name.") % (self.parent, self.columns[1], + self.columns[0], self.key)) def copy(self): return ColumnProperty(deferred=self.deferred, group=self.group, *self.columns) @@ -98,6 +113,7 @@ class ColumnProperty(StrategizedProperty): col = self.__clause_element__() return op(col._bind_param(other), col, **kwargs) + # TODO: legacy..do we need this ? (0.5) ColumnComparator = Comparator def __str__(self): @@ -117,13 +133,14 @@ class CompositeProperty(ColumnProperty): self.composite_class = class_ self.strategy_class = strategies.CompositeColumnLoader - def do_init(self): - super(ColumnProperty, self).do_init() - # TODO: similar PK check as ColumnProperty does ? - def copy(self): return CompositeProperty(deferred=self.deferred, group=self.group, composite_class=self.composite_class, *self.columns) + def do_init(self): + # skip over ColumnProperty's do_init(), + # which issues assertions that do not apply to CompositeColumnProperty + super(ColumnProperty, self).do_init() + def getattr(self, state, column): obj = state.get_impl(self.key).get(state) return self.get_col_value(column, obj) @@ -176,6 +193,47 @@ class CompositeProperty(ColumnProperty): def __str__(self): return str(self.parent.class_.__name__) + "." + self.key +class ConcreteInheritedProperty(MapperProperty): + extension = None + + def setup(self, context, entity, path, adapter, **kwargs): + pass + + def create_row_processor(self, selectcontext, path, mapper, row, adapter): + return (None, None) + + def instrument_class(self, mapper): + def warn(): + raise AttributeError("Concrete %s does not implement attribute %r at " + "the instance level. Add this property explicitly to %s." % + (self.parent, self.key, self.parent)) + + class NoninheritedConcreteProp(object): + def __set__(s, obj, value): + warn() + def __delete__(s, obj): + warn() + def __get__(s, obj, owner): + warn() + + comparator_callable = None + # TODO: put this process into a deferred callable? + for m in self.parent.iterate_to_root(): + p = m._get_property(self.key) + if not isinstance(p, ConcreteInheritedProperty): + comparator_callable = p.comparator_factory + break + + attributes.register_descriptor( + mapper.class_, + self.key, + comparator=comparator_callable(self, mapper), + parententity=mapper, + property_=self, + proxy_property=NoninheritedConcreteProp() + ) + + class SynonymProperty(MapperProperty): extension = None @@ -193,10 +251,9 @@ class SynonymProperty(MapperProperty): def create_row_processor(self, selectcontext, path, mapper, row, adapter): return (None, None) - def do_init(self): + def instrument_class(self, mapper): class_ = self.parent.class_ - self.logger.info("register managed attribute %s on class %s" % (self.key, class_.__name__)) if self.descriptor is None: class SynonymProp(object): def __set__(s, obj, value): @@ -219,8 +276,14 @@ class SynonymProperty(MapperProperty): return prop.comparator_factory(prop, mapper) return comparator - strategies.DefaultColumnLoader(self)._register_attribute( - None, None, False, comparator_callable, proxy_property=self.descriptor) + attributes.register_descriptor( + mapper.class_, + self.key, + comparator=comparator_callable(self, mapper), + parententity=mapper, + property_=self, + proxy_property=self.descriptor + ) def merge(self, session, source, dest, dont_load, _recursive): pass @@ -237,10 +300,17 @@ class ComparableProperty(MapperProperty): self.comparator_factory = comparator_factory util.set_creation_order(self) - def do_init(self): + def instrument_class(self, mapper): """Set up a proxy to the unmanaged descriptor.""" - strategies.DefaultColumnLoader(self)._register_attribute(None, None, False, self.comparator_factory, proxy_property=self.descriptor) + attributes.register_descriptor( + mapper.class_, + self.key, + comparator=self.comparator_factory(self, mapper), + parententity=mapper, + property_=self, + proxy_property=self.descriptor + ) def setup(self, context, entity, path, adapter, **kwargs): pass @@ -258,15 +328,22 @@ class RelationProperty(StrategizedProperty): """ def __init__(self, argument, - secondary=None, primaryjoin=None, secondaryjoin=None, - foreign_keys=None, uselist=None, order_by=False, backref=None, - _is_backref=False, post_update=False, cascade=False, - extension=None, viewonly=False, lazy=True, - collection_class=None, passive_deletes=False, - passive_updates=True, remote_side=None, - enable_typechecks=True, join_depth=None, - comparator_factory=None, strategy_class=None, - _local_remote_pairs=None, query_class=None): + secondary=None, primaryjoin=None, + secondaryjoin=None, + foreign_keys=None, + uselist=None, + order_by=False, + backref=None, + back_populates=None, + post_update=False, + cascade=False, extension=None, + viewonly=False, lazy=True, + collection_class=None, passive_deletes=False, + passive_updates=True, remote_side=None, + enable_typechecks=True, join_depth=None, + comparator_factory=None, + strategy_class=None, _local_remote_pairs=None, query_class=None): + self.uselist = uselist self.argument = argument self.secondary = secondary @@ -304,7 +381,7 @@ class RelationProperty(StrategizedProperty): else: self.strategy_class = strategies.LazyLoader - self._reverse_property = None + self._reverse_property = set() if cascade is not False: self.cascade = CascadeOptions(cascade) @@ -316,21 +393,37 @@ class RelationProperty(StrategizedProperty): self.order_by = order_by - if isinstance(backref, str): + self.back_populates = back_populates + + if self.back_populates: + if backref: + raise sa_exc.ArgumentError("backref and back_populates keyword arguments are mutually exclusive") + self.backref = None + elif isinstance(backref, str): # propagate explicitly sent primary/secondary join conditions to the BackRef object if # just a string was sent if secondary is not None: # reverse primary/secondary in case of a many-to-many - self.backref = BackRef(backref, primaryjoin=secondaryjoin, secondaryjoin=primaryjoin, passive_updates=self.passive_updates) + self.backref = BackRef(backref, primaryjoin=secondaryjoin, + secondaryjoin=primaryjoin, passive_updates=self.passive_updates) else: - self.backref = BackRef(backref, primaryjoin=primaryjoin, secondaryjoin=secondaryjoin, passive_updates=self.passive_updates) + self.backref = BackRef(backref, primaryjoin=primaryjoin, + secondaryjoin=secondaryjoin, passive_updates=self.passive_updates) else: self.backref = backref - self._is_backref = _is_backref + + def instrument_class(self, mapper): + attributes.register_descriptor( + mapper.class_, + self.key, + comparator=self.comparator_factory(self, mapper), + parententity=mapper, + property_=self + ) class Comparator(PropComparator): def __init__(self, prop, mapper, of_type=None, adapter=None): - self.prop = self.property = prop + self.prop = prop self.mapper = mapper self.adapter = adapter if of_type: @@ -341,14 +434,14 @@ class RelationProperty(StrategizedProperty): on the local side of generated expressions. """ - return self.__class__(self.prop, self.mapper, getattr(self, '_of_type', None), adapter) - + return self.__class__(self.property, self.mapper, getattr(self, '_of_type', None), adapter) + @property def parententity(self): - return self.prop.parent + return self.property.parent def __clause_element__(self): - elem = self.prop.parent._with_polymorphic_selectable + elem = self.property.parent._with_polymorphic_selectable if self.adapter: return self.adapter(elem) else: @@ -361,7 +454,7 @@ class RelationProperty(StrategizedProperty): return op(self, *other, **kwargs) def of_type(self, cls): - return RelationProperty.Comparator(self.prop, self.mapper, cls) + return RelationProperty.Comparator(self.property, self.mapper, cls) def in_(self, other): raise NotImplementedError("in_() not yet supported for relations. For a " @@ -371,20 +464,20 @@ class RelationProperty(StrategizedProperty): def __eq__(self, other): if other is None: - if self.prop.direction in [ONETOMANY, MANYTOMANY]: + if self.property.direction in [ONETOMANY, MANYTOMANY]: return ~self._criterion_exists() else: - return self.prop._optimized_compare(None, adapt_source=self.adapter) - elif self.prop.uselist: + return self.property._optimized_compare(None, adapt_source=self.adapter) + elif self.property.uselist: raise sa_exc.InvalidRequestError("Can't compare a collection to an object or collection; use contains() to test for membership.") else: - return self.prop._optimized_compare(other, adapt_source=self.adapter) + return self.property._optimized_compare(other, adapt_source=self.adapter) def _criterion_exists(self, criterion=None, **kwargs): if getattr(self, '_of_type', None): target_mapper = self._of_type to_selectable = target_mapper._with_polymorphic_selectable - if self.prop._is_self_referential(): + if self.property._is_self_referential(): to_selectable = to_selectable.alias() single_crit = target_mapper._single_table_criterion @@ -402,10 +495,10 @@ class RelationProperty(StrategizedProperty): source_selectable = None pj, sj, source, dest, secondary, target_adapter = \ - self.prop._create_joins(dest_polymorphic=True, dest_selectable=to_selectable, source_selectable=source_selectable) + self.property._create_joins(dest_polymorphic=True, dest_selectable=to_selectable, source_selectable=source_selectable) for k in kwargs: - crit = self.prop.mapper.class_manager[k] == kwargs[k] + crit = self.property.mapper.class_manager[k] == kwargs[k] if criterion is None: criterion = crit else: @@ -417,7 +510,7 @@ class RelationProperty(StrategizedProperty): if sj: j = _orm_annotate(pj) & sj else: - j = _orm_annotate(pj, exclude=self.prop.remote_side) + j = _orm_annotate(pj, exclude=self.property.remote_side) if criterion and target_adapter: # limit this adapter to annotated only? @@ -434,34 +527,34 @@ class RelationProperty(StrategizedProperty): return sql.exists([1], crit, from_obj=dest).correlate(source) def any(self, criterion=None, **kwargs): - if not self.prop.uselist: + if not self.property.uselist: raise sa_exc.InvalidRequestError("'any()' not implemented for scalar attributes. Use has().") return self._criterion_exists(criterion, **kwargs) def has(self, criterion=None, **kwargs): - if self.prop.uselist: + if self.property.uselist: raise sa_exc.InvalidRequestError("'has()' not implemented for collections. Use any().") return self._criterion_exists(criterion, **kwargs) def contains(self, other, **kwargs): - if not self.prop.uselist: + if not self.property.uselist: raise sa_exc.InvalidRequestError("'contains' not implemented for scalar attributes. Use ==") - clause = self.prop._optimized_compare(other, adapt_source=self.adapter) + clause = self.property._optimized_compare(other, adapt_source=self.adapter) - if self.prop.secondaryjoin: + if self.property.secondaryjoin: clause.negation_clause = self.__negated_contains_or_equals(other) return clause def __negated_contains_or_equals(self, other): - if self.prop.direction == MANYTOONE: + if self.property.direction == MANYTOONE: state = attributes.instance_state(other) - strategy = self.prop._get_strategy(strategies.LazyLoader) + strategy = self.property._get_strategy(strategies.LazyLoader) def state_bindparam(state, col): o = state.obj() # strong ref - return lambda: self.prop.mapper._get_committed_attr_by_column(o, col) + return lambda: self.property.mapper._get_committed_attr_by_column(o, col) def adapt(col): if self.adapter: @@ -474,22 +567,27 @@ class RelationProperty(StrategizedProperty): sql.or_( adapt(x) != state_bindparam(state, y), adapt(x) == None) - for (x, y) in self.prop.local_remote_pairs]) + for (x, y) in self.property.local_remote_pairs]) - criterion = sql.and_(*[x==y for (x, y) in zip(self.prop.mapper.primary_key, self.prop.mapper.primary_key_from_instance(other))]) + criterion = sql.and_(*[x==y for (x, y) in zip(self.property.mapper.primary_key, self.property.mapper.primary_key_from_instance(other))]) return ~self._criterion_exists(criterion) def __ne__(self, other): if other is None: - if self.prop.direction == MANYTOONE: - return sql.or_(*[x!=None for x in self.prop._foreign_keys]) + if self.property.direction == MANYTOONE: + return sql.or_(*[x!=None for x in self.property._foreign_keys]) else: return self._criterion_exists() - elif self.prop.uselist: + elif self.property.uselist: raise sa_exc.InvalidRequestError("Can't compare a collection to an object or collection; use contains() to test for membership.") else: return self.__negated_contains_or_equals(other) + @util.memoized_property + def property(self): + self.prop.parent.compile() + return self.prop + def compare(self, op, value, value_is_parent=False): if op == operators.eq: if value is None: @@ -512,8 +610,11 @@ class RelationProperty(StrategizedProperty): return str(self.parent.class_.__name__) + "." + self.key def merge(self, session, source, dest, dont_load, _recursive): - if not dont_load and self._reverse_property and (source, self._reverse_property) in _recursive: - return + if not dont_load: + # TODO: no test coverage for recursive check + for r in self._reverse_property: + if (source, r) in _recursive: + return source_state = attributes.instance_state(source) dest_state = attributes.instance_state(dest) @@ -547,7 +648,7 @@ class RelationProperty(StrategizedProperty): obj = session.merge(current, dont_load=dont_load, _recursive=_recursive) if obj is not None: if dont_load: - dest.__dict__[self.key] = obj + dest_state.dict[self.key] = obj else: setattr(dest, self.key, obj) @@ -567,42 +668,47 @@ class RelationProperty(StrategizedProperty): for c in instances: if c is not None and c not in visited_instances and (halt_on is None or not halt_on(c)): if not isinstance(c, self.mapper.class_): - raise AssertionError("Attribute '%s' on class '%s' doesn't handle objects of type '%s'" % (self.key, str(self.parent.class_), str(c.__class__))) + raise AssertionError("Attribute '%s' on class '%s' doesn't handle objects " + "of type '%s'" % (self.key, str(self.parent.class_), str(c.__class__))) visited_instances.add(c) # cascade using the mapper local to this object, so that its individual properties are located instance_mapper = object_mapper(c) yield (c, instance_mapper, attributes.instance_state(c)) - def _get_target_class(self): - """Return the target class of the relation, even if the - property has not been initialized yet. - - """ - if isinstance(self.argument, type): - return self.argument - else: - return self.argument.class_ - + def _add_reverse_property(self, key): + other = self.mapper._get_property(key) + self._reverse_property.add(other) + other._reverse_property.add(self) + + if not other._get_target().common_parent(self.parent): + raise sa_exc.ArgumentError("reverse_property %r on relation %s references " + "relation %s, which does not reference mapper %s" % (key, self, other, self.parent)) + def do_init(self): - self._determine_targets() + self._get_target() + self._process_dependent_arguments() self._determine_joins() self._determine_synchronize_pairs() self._determine_direction() self._determine_local_remote_pairs() self._post_init() - def _determine_targets(self): - if isinstance(self.argument, type): - self.mapper = mapper.class_mapper(self.argument, compile=False) - elif isinstance(self.argument, mapper.Mapper): - self.mapper = self.argument - elif util.callable(self.argument): - # accept a callable to suit various deferred-configurational schemes - self.mapper = mapper.class_mapper(self.argument(), compile=False) - else: - raise sa_exc.ArgumentError("relation '%s' expects a class or a mapper argument (received: %s)" % (self.key, type(self.argument))) - assert isinstance(self.mapper, mapper.Mapper), self.mapper + def _get_target(self): + if not hasattr(self, 'mapper'): + if isinstance(self.argument, type): + self.mapper = mapper.class_mapper(self.argument, compile=False) + elif isinstance(self.argument, mapper.Mapper): + self.mapper = self.argument + elif util.callable(self.argument): + # accept a callable to suit various deferred-configurational schemes + self.mapper = mapper.class_mapper(self.argument(), compile=False) + else: + raise sa_exc.ArgumentError("relation '%s' expects a class or a mapper argument (received: %s)" % (self.key, type(self.argument))) + assert isinstance(self.mapper, mapper.Mapper), self.mapper + return self.mapper + + def _process_dependent_arguments(self): # accept callables for other attributes which may require deferred initialization for attr in ('order_by', 'primaryjoin', 'secondaryjoin', 'secondary', '_foreign_keys', 'remote_side'): @@ -855,6 +961,11 @@ class RelationProperty(StrategizedProperty): # primary property handler, set up class attributes if self.is_primary(): + if self.back_populates: + self.extension = util.to_list(self.extension) or [] + self.extension.append(attributes.GenericBackrefExtension(self.back_populates)) + self._add_reverse_property(self.back_populates) + if self.backref is not None: self.backref.compile(self) elif not mapper.class_mapper(self.parent.class_, compile=False)._get_property(self.key, raiseerr=False): @@ -862,7 +973,7 @@ class RelationProperty(StrategizedProperty): "a non-primary mapper on class '%s'. New relations can only be " "added to the primary mapper, i.e. the very first " "mapper created for class '%s' " % (self.key, self.parent.class_.__name__, self.parent.class_.__name__)) - + super(RelationProperty, self).do_init() def _refers_to_parent_table(self): @@ -973,7 +1084,12 @@ log.class_logger(RelationProperty) class BackRef(object): """Attached to a RelationProperty to indicate a complementary reverse relationship. - Can optionally create the complementing RelationProperty if one does not exist already.""" + Handles the job of creating the opposite RelationProperty according to configuration. + + Alternatively, two explicit RelationProperty objects can be associated bidirectionally + using the back_populates keyword argument on each. + + """ def __init__(self, key, _prop=None, **kwargs): self.key = key @@ -1006,13 +1122,11 @@ class BackRef(object): relation = RelationProperty(parent, prop.secondary, pj, sj, backref=BackRef(prop.key, _prop=prop), - _is_backref=True, **self.kwargs) - mapper._compile_property(self.key, relation); + mapper._configure_property(self.key, relation); - prop._reverse_property = mapper._get_property(self.key) - mapper._get_property(self.key)._reverse_property = prop + prop._add_reverse_property(self.key) else: raise sa_exc.ArgumentError("Error creating backref '%s' on relation '%s': " @@ -1021,3 +1135,5 @@ class BackRef(object): mapper.ColumnProperty = ColumnProperty mapper.SynonymProperty = SynonymProperty mapper.ComparableProperty = ComparableProperty +mapper.RelationProperty = RelationProperty +mapper.ConcreteInheritedProperty = ConcreteInheritedProperty
\ No newline at end of file diff --git a/lib/sqlalchemy/orm/query.py b/lib/sqlalchemy/orm/query.py index ff7c74532..f225f346a 100644 --- a/lib/sqlalchemy/orm/query.py +++ b/lib/sqlalchemy/orm/query.py @@ -1822,8 +1822,8 @@ class _ColumnEntity(_QueryEntity): if isinstance(column, basestring): column = sql.literal_column(column) self._result_label = column.name - elif isinstance(column, (attributes.QueryableAttribute, mapper.Mapper._CompileOnAttr)): - self._result_label = column.impl.key + elif isinstance(column, attributes.QueryableAttribute): + self._result_label = column.property.key column = column.__clause_element__() else: self._result_label = getattr(column, 'key', None) diff --git a/lib/sqlalchemy/orm/session.py b/lib/sqlalchemy/orm/session.py index b44ec25d5..cb79d7dc2 100644 --- a/lib/sqlalchemy/orm/session.py +++ b/lib/sqlalchemy/orm/session.py @@ -1521,7 +1521,7 @@ class Session(object): return util.IdentitySet(self._new.values()) _expire_state = attributes.InstanceState.expire_attributes -register_attribute = unitofwork.register_attribute +UOWEventHandler = unitofwork.UOWEventHandler _sessions = weakref.WeakValueDictionary() diff --git a/lib/sqlalchemy/orm/strategies.py b/lib/sqlalchemy/orm/strategies.py index a159e4bfa..a4c1f7d2d 100644 --- a/lib/sqlalchemy/orm/strategies.py +++ b/lib/sqlalchemy/orm/strategies.py @@ -18,37 +18,71 @@ from sqlalchemy.orm.interfaces import ( from sqlalchemy.orm import session as sessionlib from sqlalchemy.orm import util as mapperutil +def _register_attribute(strategy, useobject, + compare_function=None, + typecallable=None, + copy_function=None, + mutable_scalars=False, + uselist=False, + callable_=None, + proxy_property=None, + active_history=False, + impl_class=None, + **kw +): + + prop = strategy.parent_property + attribute_ext = util.to_list(prop.extension) or [] + if getattr(prop, 'backref', None): + attribute_ext.append(prop.backref.extension) + + if prop.key in prop.parent._validators: + attribute_ext.append(mapperutil.Validator(prop.key, prop.parent._validators[prop.key])) + + if useobject: + attribute_ext.append(sessionlib.UOWEventHandler(prop.key)) + + for mapper in prop.parent.polymorphic_iterator(): + if (mapper is prop.parent or not mapper.concrete) and mapper.has_property(prop.key): + attributes.register_attribute_impl( + mapper.class_, + prop.key, + parent_token=prop, + mutable_scalars=mutable_scalars, + uselist=uselist, + copy_function=copy_function, + compare_function=compare_function, + useobject=useobject, + extension=attribute_ext, + trackparent=useobject, + typecallable=typecallable, + callable_=callable_, + active_history=active_history, + impl_class=impl_class, + **kw + ) -class DefaultColumnLoader(LoaderStrategy): - def _register_attribute(self, compare_function, copy_function, mutable_scalars, - comparator_factory, callable_=None, proxy_property=None, active_history=False): - self.logger.info("%s register managed attribute" % self) - - attribute_ext = util.to_list(self.parent_property.extension) or [] - if self.key in self.parent._validators: - attribute_ext.append(mapperutil.Validator(self.key, self.parent._validators[self.key])) - - for mapper in self.parent.polymorphic_iterator(): - if (mapper is self.parent or not mapper.concrete) and mapper.has_property(self.key): - sessionlib.register_attribute( - mapper.class_, - self.key, - uselist=False, - useobject=False, - copy_function=copy_function, - compare_function=compare_function, - mutable_scalars=mutable_scalars, - comparator=comparator_factory(self.parent_property, mapper), - parententity=mapper, - callable_=callable_, - extension=attribute_ext, - proxy_property=proxy_property, - active_history=active_history - ) - -log.class_logger(DefaultColumnLoader) +class UninstrumentedColumnLoader(LoaderStrategy): + """Represent the strategy for a MapperProperty that doesn't instrument the class. -class ColumnLoader(DefaultColumnLoader): + The polymorphic_on argument of mapper() often results in this, + if the argument is against the with_polymorphic selectable. + + """ + def init(self): + self.columns = self.parent_property.columns + + def setup_query(self, context, entity, path, adapter, column_collection=None, **kwargs): + for c in self.columns: + if adapter: + c = adapter.columns[c] + column_collection.append(c) + + def create_row_processor(self, selectcontext, path, mapper, row, adapter): + return (None, None) + +class ColumnLoader(LoaderStrategy): + """Strategize the loading of a plain column-based MapperProperty.""" def init(self): self.columns = self.parent_property.columns @@ -64,14 +98,12 @@ class ColumnLoader(DefaultColumnLoader): self.is_class_level = True coltype = self.columns[0].type active_history = self.columns[0].primary_key # TODO: check all columns ? check for foreign Key as well? - - self._register_attribute( - coltype.compare_values, - coltype.copy_value, - self.columns[0].type.is_mutable(), - self.parent_property.comparator_factory, + + _register_attribute(self, useobject=False, + compare_function=coltype.compare_values, + copy_function=coltype.copy_value, + mutable_scalars=self.columns[0].type.is_mutable(), active_history = active_history - ) def create_row_processor(self, selectcontext, path, mapper, row, adapter): @@ -99,6 +131,8 @@ class ColumnLoader(DefaultColumnLoader): log.class_logger(ColumnLoader) class CompositeColumnLoader(ColumnLoader): + """Strategize the loading of a composite column-based MapperProperty.""" + def init_class_attribute(self): self.is_class_level = True self.logger.info("%s register managed composite attribute" % self) @@ -120,11 +154,11 @@ class CompositeColumnLoader(ColumnLoader): else: return True - self._register_attribute( - compare, - copy, - True, - self.parent_property.comparator_factory + _register_attribute(self, useobject=False, + compare_function=compare, + copy_function=copy, + mutable_scalars=True + #active_history ? ) def create_row_processor(self, selectcontext, path, mapper, row, adapter): @@ -153,8 +187,8 @@ class CompositeColumnLoader(ColumnLoader): log.class_logger(CompositeColumnLoader) -class DeferredColumnLoader(DefaultColumnLoader): - """Deferred column loader, a per-column or per-column-group lazy loader.""" +class DeferredColumnLoader(LoaderStrategy): + """Strategize the loading of a deferred column-based MapperProperty.""" def create_row_processor(self, selectcontext, path, mapper, row, adapter): col = self.columns[0] @@ -184,11 +218,11 @@ class DeferredColumnLoader(DefaultColumnLoader): def init_class_attribute(self): self.is_class_level = True - self._register_attribute( - self.columns[0].type.compare_values, - self.columns[0].type.copy_value, - self.columns[0].type.is_mutable(), - self.parent_property.comparator_factory, + + _register_attribute(self, useobject=False, + compare_function=self.columns[0].type.compare_values, + copy_function=self.columns[0].type.copy_value, + mutable_scalars=self.columns[0].type.is_mutable(), callable_=self.class_level_loader, ) @@ -282,6 +316,8 @@ class UndeferGroupOption(MapperOption): query._attributes[('undefer', self.group)] = True class AbstractRelationLoader(LoaderStrategy): + """LoaderStratgies which deal with related objects as opposed to scalars.""" + def init(self): for attr in ['mapper', 'target', 'table', 'uselist']: setattr(self, attr, getattr(self.parent_property, attr)) @@ -291,37 +327,18 @@ class AbstractRelationLoader(LoaderStrategy): state.set_callable(self.key, callable_) else: state.initialize(self.key) - - def _register_attribute(self, class_, callable_=None, impl_class=None, **kwargs): - self.logger.info("%s register managed %s attribute" % (self, (self.uselist and "collection" or "scalar"))) - - attribute_ext = util.to_list(self.parent_property.extension) or [] - - if self.parent_property.backref: - attribute_ext.append(self.parent_property.backref.extension) - - if self.key in self.parent._validators: - attribute_ext.append(mapperutil.Validator(self.key, self.parent._validators[self.key])) - - sessionlib.register_attribute( - class_, - self.key, - uselist=self.uselist, - useobject=True, - extension=attribute_ext, - trackparent=True, - typecallable=self.parent_property.collection_class, - callable_=callable_, - comparator=self.parent_property.comparator, - parententity=self.parent, - impl_class=impl_class, - **kwargs - ) class NoLoader(AbstractRelationLoader): + """Strategize a relation() that doesn't load data automatically.""" + def init_class_attribute(self): self.is_class_level = True - self._register_attribute(self.parent.class_) + + _register_attribute(self, + useobject=True, + uselist=self.parent_property.uselist, + typecallable = self.parent_property.collection_class, + ) def create_row_processor(self, selectcontext, path, mapper, row, adapter): def new_execute(state, row, **flags): @@ -336,6 +353,8 @@ class NoLoader(AbstractRelationLoader): log.class_logger(NoLoader) class LazyLoader(AbstractRelationLoader): + """Strategize a relation() that loads when first accessed.""" + def init(self): super(LazyLoader, self).init() (self.__lazywhere, self.__bind_to_col, self._equated_columns) = self._create_lazy_clause(self.parent_property) @@ -351,7 +370,14 @@ class LazyLoader(AbstractRelationLoader): def init_class_attribute(self): self.is_class_level = True - self._register_attribute(self.parent.class_, callable_=self.class_level_loader) + + + _register_attribute(self, + useobject=True, + callable_=self.class_level_loader, + uselist = self.parent_property.uselist, + typecallable = self.parent_property.collection_class, + ) def lazy_clause(self, state, reverse_direction=False, alias_secondary=False, adapt_source=None): if state is None: @@ -564,7 +590,7 @@ class LoadLazyAttribute(object): return None class EagerLoader(AbstractRelationLoader): - """Loads related objects inline with a parent query.""" + """Strategize a relation() that loads within the process of the parent object being selected.""" def init(self): super(EagerLoader, self).init() diff --git a/lib/sqlalchemy/orm/unitofwork.py b/lib/sqlalchemy/orm/unitofwork.py index 4efab88ae..0c282a7b8 100644 --- a/lib/sqlalchemy/orm/unitofwork.py +++ b/lib/sqlalchemy/orm/unitofwork.py @@ -69,23 +69,6 @@ class UOWEventHandler(interfaces.AttributeExtension): sess.expunge(oldvalue) return newvalue -def register_attribute(class_, key, *args, **kwargs): - """Register an attribute with the attributes module. - - Overrides attributes.register_attribute() to add - unitofwork-specific event handlers. - - """ - useobject = kwargs.get('useobject', False) - if useobject: - # for object-holding attributes, instrument UOWEventHandler - # to process per-attribute cascades - extension = util.to_list(kwargs.pop('extension', None) or []) - extension.append(UOWEventHandler(key)) - - kwargs['extension'] = extension - return attributes.register_attribute(class_, key, *args, **kwargs) - class UOWTransaction(object): """Handles the details of organizing and executing transaction diff --git a/lib/sqlalchemy/orm/util.py b/lib/sqlalchemy/orm/util.py index 4f99586da..9abcf90dd 100644 --- a/lib/sqlalchemy/orm/util.py +++ b/lib/sqlalchemy/orm/util.py @@ -304,8 +304,8 @@ class AliasedClass(object): existing = getattr(self.__target, prop.key) comparator = existing.comparator.adapted(self.__adapt_element) - queryattr = attributes.QueryableAttribute( - existing.impl, parententity=self, comparator=comparator) + queryattr = attributes.QueryableAttribute(prop.key, + impl=existing.impl, parententity=self, comparator=comparator) setattr(self, prop.key, queryattr) return queryattr diff --git a/lib/sqlalchemy/schema.py b/lib/sqlalchemy/schema.py index 32ea2b5ee..734614807 100644 --- a/lib/sqlalchemy/schema.py +++ b/lib/sqlalchemy/schema.py @@ -1431,7 +1431,7 @@ class Index(SchemaItem): def _init_items(self, *args): for column in args: - self.append_column(column) + self.append_column(_to_schema_column(column)) def _set_parent(self, table): self.table = table @@ -2107,7 +2107,13 @@ class DDL(object): for key in ('on', 'context') if getattr(self, key)])) - +def _to_schema_column(element): + if hasattr(element, '__clause_element__'): + element = element.__clause_element__() + if not isinstance(element, Column): + raise exc.ArgumentError("schema.Column object expected") + return element + def _bind_or_error(schemaitem): bind = schemaitem.bind if not bind: |
