Author: ianb
Date: 2005-11-02 04:46:27 +0000 (Wed, 02 Nov 2005)
New Revision: 1185
Added:
SQLObject/trunk/sqlobject/tests/test_new_joins.py
Modified:
SQLObject/trunk/sqlobject/declarative.py
SQLObject/trunk/sqlobject/joins.py
SQLObject/trunk/sqlobject/main.py
Log:
New joins! ManyToMany and OneToMany; not fully documented yet, but still more sensible and smarter. Also (in main.py) call all the events promised in events.py.
Modified: SQLObject/trunk/sqlobject/declarative.py
===================================================================
--- SQLObject/trunk/sqlobject/declarative.py 2005-11-02 04:45:26 UTC (rev 1184)
+++ SQLObject/trunk/sqlobject/declarative.py 2005-11-02 04:46:27 UTC (rev 1185)
@@ -33,6 +33,7 @@
"""
from __future__ import generators
+import events
__all__ = ('classinstancemethod', 'DeclarativeMeta', 'Declarative')
@@ -88,10 +89,16 @@
class DeclarativeMeta(type):
def __new__(meta, class_name, bases, new_attrs):
+ post_funcs = []
+ events.send(events.ClassCreateSignal,
+ bases[0], class_name, bases, new_attrs,
+ post_funcs)
cls = type.__new__(meta, class_name, bases, new_attrs)
if new_attrs.has_key('__classinit__'):
cls.__classinit__ = staticmethod(cls.__classinit__.im_func)
cls.__classinit__(cls, new_attrs)
+ for func in post_funcs:
+ func(cls)
return cls
class Declarative(object):
@@ -190,3 +197,8 @@
__repr__ = classinstancemethod(__repr__)
+def setup_attributes(cls, new_attrs):
+ for name, value in new_attrs.items():
+ if hasattr(value, '__addtoclass__'):
+ value.__addtoclass__(cls, name)
+
Modified: SQLObject/trunk/sqlobject/joins.py
===================================================================
--- SQLObject/trunk/sqlobject/joins.py 2005-11-02 04:45:26 UTC (rev 1184)
+++ SQLObject/trunk/sqlobject/joins.py 2005-11-02 04:46:27 UTC (rev 1185)
@@ -3,8 +3,10 @@
import styles
import classregistry
from col import popKey
+import events
-__all__ = ['MultipleJoin', 'SQLMultipleJoin', 'RelatedJoin', 'SQLRelatedJoin', 'SingleJoin']
+__all__ = ['MultipleJoin', 'SQLMultipleJoin', 'RelatedJoin', 'SQLRelatedJoin',
+ 'SingleJoin', 'ManyToMany', 'OneToMany']
def getID(obj):
try:
@@ -263,3 +265,209 @@
class SingleJoin(Join):
baseClass = SOSingleJoin
+
+
+
+
+class ManyToMany(object):
+
+ def __init__(self, otherClassName,
+ intermediateTable=None,
+ joinColumn=None,
+ otherColumn=None,
+ createJoinTable=True):
+ self.otherClassName = otherClassName
+ self.intermediateTable = intermediateTable
+ self.joinColumn = joinColumn
+ self.otherColumn = otherColumn
+ self.createJoinTable = createJoinTable
+
+ def __addtoclass__(self, soClass, name):
+ setattr(soClass, name,
+ SOManyToMany(soClass, name=name,
+ otherClassName=self.otherClassName,
+ intermediateTable=self.intermediateTable,
+ joinColumn=self.joinColumn,
+ otherColumn=self.otherColumn,
+ createJoinTable=self.createJoinTable))
+
+class SOManyToMany(object):
+
+ def __init__(self, soClass, name, otherClassName,
+ intermediateTable, joinColumn, otherColumn,
+ createJoinTable):
+ self.name = name
+ self.intermediateTable = intermediateTable
+ self.joinColumn = joinColumn
+ self.otherColumn = otherColumn
+ self.createJoinTable = createJoinTable
+ self.soClass = self.otherClass = None
+ classregistry.registry(
+ soClass.sqlmeta.registry).addClassCallback(
+ otherClassName, self._setOtherClass)
+ classregistry.registry(
+ soClass.sqlmeta.registry).addClassCallback(
+ soClass.__name__, self._setThisClass)
+
+ def _setThisClass(self, soClass):
+ self.soClass = soClass
+ if self.soClass and self.otherClass:
+ self._finishSet()
+
+ def _setOtherClass(self, otherClass):
+ self.otherClass = otherClass
+ if self.soClass and self.otherClass:
+ self._finishSet()
+
+ def _finishSet(self):
+ if self.intermediateTable is None:
+ names = [self.soClass.sqlmeta.table,
+ self.otherClass.sqlmeta.table]
+ names.sort()
+ self.intermediateTable = '%s_%s' % (names[0], names[1])
+ if not self.otherColumn:
+ self.otherColumn = self.soClass.sqlmeta.style.tableReference(
+ self.otherClass.sqlmeta.table)
+ if not self.joinColumn:
+ self.joinColumn = styles.getStyle(
+ self.soClass).tableReference(self.soClass.sqlmeta.table)
+ events.listen(self.event_CreateTableSignal,
+ self.soClass, events.CreateTableSignal)
+ events.listen(self.event_CreateTableSignal,
+ self.otherClass, events.CreateTableSignal)
+
+ def __get__(self, obj, type):
+ if obj is None:
+ return self
+ query = (
+ (self.otherClass.q.id ==
+ sqlbuilder.Field(self.intermediateTable, self.otherColumn))
+ & (sqlbuilder.Field(self.intermediateTable, self.joinColumn)
+ == obj.id))
+ select = self.otherClass.select(query)
+ return _ManyToManySelectWrapper(obj, self, select)
+
+ def __sqlrepr__(self, dbname):
+ return self.query.__sqlrepr__(self, dbname)
+
+ def event_CreateTableSignal(self, soClass, connection, extra_sql,
+ post_funcs):
+ if self.createJoinTable:
+ post_funcs.append(self.event_CreateTableSignalPost)
+
+ def event_CreateTableSignalPost(self, soClass, connection):
+ if connection.tableExists(self.intermediateTable):
+ return
+ connection._SO_createJoinTable(self)
+
+class _ManyToManySelectWrapper(object):
+
+ def __init__(self, forObject, join, select):
+ self.forObject = forObject
+ self.join = join
+ self.select = select
+
+ def __getattr__(self, attr):
+ # @@: This passes through private variable access too... should it?
+ # Also magic methods, like __str__
+ return getattr(self, select, attr)
+
+ def __repr__(self):
+ return '<%s for: %s>' % (self.__class__.__name__, repr(self.select))
+
+ def __str__(self):
+ return str(self.select)
+
+ def __iter__(self):
+ return iter(self.select)
+
+ def __getitem__(self, key):
+ return self.select[key]
+
+ def add(self, obj):
+ print "Add", obj, "to", self.forObject
+ obj._connection._SO_intermediateInsert(
+ self.join.intermediateTable,
+ self.join.joinColumn,
+ getID(self.forObject),
+ self.join.otherColumn,
+ getID(obj))
+
+ def remove(self, obj):
+ obj._connection._SO_intermediateDelete(
+ self.join.intermediateTable,
+ self.join.joinColumn,
+ getID(self.forObject),
+ self.join.otherColumn,
+ getID(obj))
+
+ def create(self, **kw):
+ obj = self.join.otherClass(**kw)
+ self.add(obj)
+ return obj
+
+class OneToMany(object):
+
+ def __init__(self, otherClassName, joinColumn=None):
+ self.otherClassName = otherClassName
+ self.joinColumn = joinColumn
+
+ def __addtoclass__(self, soClass, name):
+ setattr(soClass, name,
+ SOOneToMany(soClass, name=name,
+ otherClassName=self.otherClassName,
+ joinColumn=self.joinColumn))
+
+class SOOneToMany(object):
+
+ def __init__(self, soClass, name, otherClassName, joinColumn):
+ self.soClass = soClass
+ self.name = name
+ self.joinColumn = joinColumn
+ classregistry.registry(
+ soClass.sqlmeta.registry).addClassCallback(
+ otherClassName, self._setOtherClass)
+
+ def _setOtherClass(self, otherClass):
+ self.otherClass = otherClass
+ if not self.joinColumn:
+ self.joinColumn = styles.getStyle(
+ self.soClass).tableReference(self.soClass.sqlmeta.table)
+
+ def __get__(self, obj, type):
+ if obj is None:
+ return self
+ query = (
+ sqlbuilder.Field(self.otherClass.sqlmeta.table, self.joinColumn)
+ == obj.id)
+ select = self.otherClass.select(query)
+ return _OneToManySelectWrapper(obj, self, select)
+
+class _OneToManySelectWrapper(object):
+
+ def __init__(self, forObject, join, select):
+ self.forObject = forObject
+ self.join = join
+ self.select = select
+
+ def __getattr__(self, attr):
+ # @@: This passes through private variable access too... should it?
+ # Also magic methods, like __str__
+ return getattr(self, select, attr)
+
+ def __repr__(self):
+ return '<%s for: %s>' % (self.__class__.__name__, repr(self.select))
+
+ def __str__(self):
+ return str(self.select)
+
+ def __iter__(self):
+ return iter(self.select)
+
+ def __getitem__(self, key):
+ return self.select[key]
+
+ def create(self, **kw):
+ kw[self.join.joinColumn] = self.forObject.id
+ return self.join.otherClass(**kw)
+
Modified: SQLObject/trunk/sqlobject/main.py
===================================================================
--- SQLObject/trunk/sqlobject/main.py 2005-11-02 04:45:26 UTC (rev 1184)
+++ SQLObject/trunk/sqlobject/main.py 2005-11-02 04:46:27 UTC (rev 1185)
@@ -45,6 +45,7 @@
import index
import classregistry
import declarative
+import events
from sresults import SelectResults
from formencode import schema, compound
@@ -236,10 +237,16 @@
for attr in cls._unshared_attributes:
if not new_attrs.has_key(attr):
setattr(cls, attr, None)
+ declarative.setup_attributes(cls, new_attrs)
def __init__(self, instance):
self.instance = instance
+ def send(cls, signal, *args, **kw):
+ events.send(signal, cls.soClass, *args, **kw)
+
+ send = classmethod(send)
+
def setClass(cls, soClass):
cls.soClass = soClass
if not cls.style:
@@ -289,6 +296,9 @@
########################################
def addColumn(cls, columnDef, changeSchema=False, connection=None):
+ post_funcs = []
+ cls.send(events.AddColumnSignal, cls.soClass, connection,
+ columnDef.name, columnDef, changeSchema, post_funcs)
sqlmeta = cls
soClass = cls.soClass
del cls
@@ -396,9 +406,6 @@
setattr(soClass, setterName(name)[:-2], setter)
sqlmeta._plainForeignSetters[name[:-2]] = 1
- # We'll need to put in a real reference at
- # some point. See needSet at the top of the
- # file for more on this.
classregistry.registry(sqlmeta.registry).addClassCallback(
column.foreignKey,
lambda foreign, me, attr: setattr(me, attr, foreign),
@@ -415,6 +422,9 @@
if soClass._SO_finishedClassCreation:
makeProperties(soClass)
+ for func in post_funcs:
+ func(soClass, column)
+
addColumn = classmethod(addColumn)
def addColumnsFromDatabase(sqlmeta, connection=None):
@@ -438,6 +448,9 @@
else:
raise IndexError(
"Column with definition %r not found" % column)
+ post_funcs = []
+ cls.send(events.DeleteColumnSignal, connection, column.name, column,
+ post_funcs)
name = column.name
del sqlmeta.columns[name]
del sqlmeta.columnDefinitions[name]
@@ -465,6 +478,9 @@
if soClass._SO_finishedClassCreation:
unmakeProperties(soClass)
+ for func in post_funcs:
+ func(soClass, column)
+
delColumn = classmethod(delColumn)
########################################
@@ -666,6 +682,8 @@
def __classinit__(cls, new_attrs):
+ declarative.setup_attributes(cls, new_attrs)
+
# This is true if we're initializing the SQLObject class,
# instead of a subclass:
is_base = cls.__bases__ == (object,)
@@ -1039,6 +1057,11 @@
# in the database, and we can't insert it until all
# the parts are set. So we just keep them in a
# dictionary until later:
+ d = {name: value}
+ self.sqlmeta.send(events.RowUpdateSignal, self, d)
+ if len(d) != 1 or name not in d:
+ return self.set(**d)
+ value = d[name]
if from_python:
dbValue = from_python(value, self._SO_validatorState)
else:
@@ -1059,6 +1082,7 @@
setattr(self, instanceName(name), value)
def set(self, **kw):
+ self.sqlmeta.send(events.RowUpdateSignal, self, kw)
# set() is used to update multiple values at once,
# potentially with one SQL statement if possible.
@@ -1179,11 +1203,15 @@
if kw.has_key('_SO_fetch_no_create'):
return
+ post_funcs = []
+ self.sqlmeta.send(events.RowCreateSignal, kw, post_funcs)
+
# Pass the connection object along if we were given one.
if kw.has_key('connection'):
self._connection = kw['connection']
self.sqlmeta._perConnection = True
del kw['connection']
+
self._SO_writeLock = threading.Lock()
if kw.has_key('id'):
@@ -1193,6 +1221,8 @@
id = None
self._create(id, **kw)
+ for func in post_funcs:
+ func(self)
def _create(self, id, **kw):
@@ -1306,9 +1336,17 @@
conn = connection or cls._connection
if ifExists and not conn.tableExists(cls.sqlmeta.table):
return
+ extra_sql = []
+ post_funcs = []
+ cls.sqlmeta.send(events.DropTableSignal, cls, connection,
+ extra_sql, post_funcs)
conn.dropTable(cls.sqlmeta.table, cascade)
if dropJoinTables:
cls.dropJoinTables(ifExists=ifExists, connection=conn)
+ for sql in extra_sql:
+ connection.query(sql)
+ for func in post_funcs:
+ func(cls, conn)
dropTable = classmethod(dropTable)
def createTable(cls, ifNotExists=False, createJoinTables=True,
@@ -1317,14 +1355,21 @@
conn = connection or cls._connection
if ifNotExists and conn.tableExists(cls.sqlmeta.table):
return
+ extra_sql = []
+ post_funcs = []
+ cls.sqlmeta.send(events.CreateTableSignal, cls, connection,
+ extra_sql, post_funcs)
constraints = conn.createTable(cls)
+ extra_sql.extend(constraints)
if createJoinTables:
cls.createJoinTables(ifNotExists=ifNotExists,
connection=conn)
if createIndexes:
cls.createIndexes(ifNotExists=ifNotExists,
connection=conn)
- return constraints
+ for func in post_funcs:
+ func(cls, conn)
+ return extra_sql
createTable = classmethod(createTable)
def createTableSQL(cls, createJoinTables=True, connection=None,
@@ -1419,6 +1464,7 @@
clearTable = classmethod(clearTable)
def destroySelf(self):
+ self.sqlmeta.send(events.RowDestroySignal, self)
# Kills this object. Kills it dead!
depends = []
klass = self.__class__
Copied: SQLObject/trunk/sqlobject/tests/test_new_joins.py (from rev 1174, SQLObject/trunk/sqlobject/tests/test_joins.py)
===================================================================
--- SQLObject/trunk/sqlobject/tests/test_joins.py 2005-10-28 16:29:54 UTC (rev 1174)
+++ SQLObject/trunk/sqlobject/tests/test_new_joins.py 2005-11-02 04:46:27 UTC (rev 1185)
@@ -0,0 +1,153 @@
+from sqlobject import *
+from sqlobject.tests.dbtest import *
+
+########################################
+## Joins
+########################################
+
+class PersonJoinerNew(SQLObject):
+
+ name = StringCol(length=40, alternateID=True)
+ addressJoiners = ManyToMany('AddressJoinerNew')
+
+class AddressJoinerNew(SQLObject):
+
+ zip = StringCol(length=5, alternateID=True)
+ personJoiners = ManyToMany('PersonJoinerNew')
+
+class ImplicitJoiningSONew(SQLObject):
+ foo = ManyToMany('Bar')
+
+class ExplicitJoiningSONew(SQLObject):
+ foo = OneToMany('Bar')
+
+class TestJoin:
+
+ def setup_method(self, meth):
+ setupClass(PersonJoinerNew)
+ setupClass(AddressJoinerNew)
+ for n in ['bob', 'tim', 'jane', 'joe', 'fred', 'barb']:
+ PersonJoinerNew(name=n)
+ for z in ['11111', '22222', '33333', '44444']:
+ AddressJoinerNew(zip=z)
+
+ def test_join(self):
+ b = PersonJoinerNew.byName('bob')
+ assert list(b.addressJoiners) == []
+ z = AddressJoinerNew.byZip('11111')
+ b.addressJoiners.add(z)
+ self.assertZipsEqual(b.addressJoiners, ['11111'])
+ print str(z.personJoiners), repr(z.personJoiners)
+ self.assertNamesEqual(z.personJoiners, ['bob'])
+ z2 = AddressJoinerNew.byZip('22222')
+ b.addressJoiners.add(z2)
+ print str(b.addressJoiners)
+ self.assertZipsEqual(b.addressJoiners, ['11111', '22222'])
+ self.assertNamesEqual(z2.personJoiners, ['bob'])
+ b.addressJoiners.remove(z)
+ self.assertZipsEqual(b.addressJoiners, ['22222'])
+ self.assertNamesEqual(z.personJoiners, [])
+
+ def assertZipsEqual(self, zips, dest):
+ assert [a.zip for a in zips] == dest
+
+ def assertNamesEqual(self, people, dest):
+ assert [p.name for p in people] == dest
+
+ def test_joinAttributeWithUnderscores(self):
+ # Make sure that the implicit setting of joinMethodName works
+ assert hasattr(ImplicitJoiningSONew, 'foo')
+ assert not hasattr(ImplicitJoiningSONew, 'bars')
+
+ # And make sure explicit setting also works
+ assert hasattr(ExplicitJoiningSONew, 'foo')
+ assert not hasattr(ExplicitJoiningSONew, 'bars')
+
+
+class PersonJoinerNew2(SQLObject):
+
+ name = StringCol('name', length=40, alternateID=True)
+ addressJoiner2s = OneToMany('AddressJoinerNew2')
+
+class AddressJoinerNew2(SQLObject):
+
+ class sqlmeta:
+ defaultOrder = ['-zip', 'plus4']
+
+ zip = StringCol(length=5)
+ plus4 = StringCol(length=4, default=None)
+ personJoinerNew2 = ForeignKey('PersonJoinerNew2')
+
+class TestJoin2:
+
+ def setup_method(self, meth):
+ setupClass([PersonJoinerNew2, AddressJoinerNew2])
+ p1 = PersonJoinerNew2(name='bob')
+ p2 = PersonJoinerNew2(name='sally')
+ for z in ['11111', '22222', '33333']:
+ a = AddressJoinerNew2(zip=z, personJoinerNew2=p1)
+ #p1.addAddressJoinerNew2(a)
+ AddressJoinerNew2(zip='00000', personJoinerNew2=p2)
+
+ def test_basic(self):
+ bob = PersonJoinerNew2.byName('bob')
+ sally = PersonJoinerNew2.byName('sally')
+ print bob.addressJoiner2s
+ print bob
+ assert len(list(bob.addressJoiner2s)) == 3
+ assert len(list(sally.addressJoiner2s)) == 1
+ bob.addressJoiner2s[0].destroySelf()
+ assert len(list(bob.addressJoiner2s)) == 2
+ z = bob.addressJoiner2s[0]
+ z.zip = 'xxxxx'
+ id = z.id
+ del z
+ z = AddressJoinerNew2.get(id)
+ assert z.zip == 'xxxxx'
+
+ def test_defaultOrder(self):
+ p1 = PersonJoinerNew2.byName('bob')
+ assert ([i.zip for i in p1.addressJoiner2s]
+ == ['33333', '22222', '11111'])
+
+
+_personJoiner3_getters = []
+_personJoiner3_setters = []
+
+class PersonJoinerNew3(SQLObject):
+
+ name = StringCol('name', length=40, alternateID=True)
+ addressJoinerNew3s = OneToMany('AddressJoinerNew3')
+
+class AddressJoinerNew3(SQLObject):
+
+ zip = StringCol(length=5)
+ personJoinerNew3 = ForeignKey('PersonJoinerNew3')
+
+ def _get_personJoinerNew3(self):
+ value = self._SO_get_personJoinerNew3()
+ _personJoiner3_getters.append((self, value))
+ return value
+
+ def _set_personJoinerNew3(self, value):
+ self._SO_set_personJoinerNew3(value)
+ _personJoiner3_setters.append((self, value))
+
+class TestJoin3:
+
+ def setup_method(self, meth):
+ setupClass([PersonJoinerNew3, AddressJoinerNew3])
+ p1 = PersonJoinerNew3(name='bob')
+ p2 = PersonJoinerNew3(name='sally')
+ for z in ['11111', '22222', '33333']:
+ a = AddressJoinerNew3(zip=z, personJoinerNew3=p1)
+ AddressJoinerNew3(zip='00000', personJoinerNew3=p2)
+
+ def test_accessors(self):
+ assert len(list(_personJoiner3_getters)) == 0
+ assert len(list(_personJoiner3_setters)) == 4
+ bob = PersonJoinerNew3.byName('bob')
+ for addressJoiner3 in bob.addressJoinerNew3s:
+ addressJoiner3.personJoinerNew3
+ assert len(list(_personJoiner3_getters)) == 3
+ assert len(list(_personJoiner3_setters)) == 4
|