Merge pull request #347 from niphlod/enhancement/DAL

first steps to a cleaner DAL and rname integration
This commit is contained in:
mdipierro
2014-01-14 20:10:26 -08:00
2 changed files with 242 additions and 224 deletions
+73 -51
View File
@@ -8611,8 +8611,7 @@ class Table(object):
db, db,
tablename, tablename,
*fields, *fields,
**args **args):
):
""" """
Initializes the table and performs checking on the provided fields. Initializes the table and performs checking on the provided fields.
@@ -8625,12 +8624,19 @@ class Table(object):
""" """
self._actual = False # set to True by define_table() self._actual = False # set to True by define_table()
self._tablename = tablename self._tablename = tablename
self._ot = None # args.get('rname') if (not isinstance(tablename, str) or tablename[0] == '_'
or hasattr(DAL, tablename) or '.' in tablename
or REGEX_PYTHON_KEYWORDS.match(tablename)
):
raise SyntaxError('Field: invalid table name: %s, '
'use rname for "funny" names' % tablename)
self._ot = None
self._rname = args.get('rname') self._rname = args.get('rname')
self._sequence_name = args.get('sequence_name') or \ self._sequence_name = (args.get('sequence_name') or
db and db._adapter.sequence_name(self._rname or tablename) db and db._adapter.sequence_name(self._rname
self._trigger_name = args.get('trigger_name') or \ or tablename))
db and db._adapter.trigger_name(tablename) self._trigger_name = (args.get('trigger_name') or
db and db._adapter.trigger_name(tablename))
self._common_filter = args.get('common_filter') self._common_filter = args.get('common_filter')
self._format = args.get('format') self._format = args.get('format')
self._singular = args.get( self._singular = args.get(
@@ -8655,10 +8661,10 @@ class Table(object):
if _primarykey is not None: if _primarykey is not None:
if not isinstance(_primarykey, list): if not isinstance(_primarykey, list):
raise SyntaxError( raise SyntaxError(
"primarykey must be a list of fields from table '%s'" \ "primarykey must be a list of fields from table '%s'"
% tablename) % tablename)
if len(_primarykey) == 1: if len(_primarykey) == 1:
self._id = [f for f in fields if isinstance(f,Field) \ self._id = [f for f in fields if isinstance(f, Field)
and f.name ==_primarykey[0]][0] and f.name ==_primarykey[0]][0]
elif not [f for f in fields if (isinstance(f, Field) and elif not [f for f in fields if (isinstance(f, Field) and
f.type == 'id') or (isinstance(f, dict) and f.type == 'id') or (isinstance(f, dict) and
@@ -8668,6 +8674,7 @@ class Table(object):
fieldnames.add('id') fieldnames.add('id')
self._id = field self._id = field
virtual_fields = [] virtual_fields = []
def include_new(field): def include_new(field):
newfields.append(field) newfields.append(field)
fieldnames.add(field.name) fieldnames.add(field.name)
@@ -8698,7 +8705,7 @@ class Table(object):
self.virtualfields = [] self.virtualfields = []
fields = list(fields) fields = list(fields)
if db and db._adapter.uploads_in_blob==True: if db and db._adapter.uploads_in_blob is True:
uploadfields = [f.name for f in fields if f.type == 'blob'] uploadfields = [f.name for f in fields if f.type == 'blob']
for field in fields: for field in fields:
fn = field.uploadfield fn = field.uploadfield
@@ -8725,8 +8732,8 @@ class Table(object):
else: else:
fname_item = field_name fname_item = field_name
if fname_item in fieldnames_set: if fname_item in fieldnames_set:
raise SyntaxError("duplicate field %s in table %s" \ raise SyntaxError("duplicate field %s in table %s" %
% (field_name, tablename)) (field_name, tablename))
else: else:
fieldnames_set.add(fname_item) fieldnames_set.add(fname_item)
@@ -8743,7 +8750,8 @@ class Table(object):
for k in _primarykey: for k in _primarykey:
if k not in self.fields: if k not in self.fields:
raise SyntaxError( raise SyntaxError(
"primarykey must be a list of fields from table '%s " % tablename) "primarykey must be a list of fields from table '%s " %
tablename)
else: else:
self[k].notnull = True self[k].notnull = True
for field in virtual_fields: for field in virtual_fields:
@@ -8773,8 +8781,9 @@ class Table(object):
clones = [] clones = []
for field in self: for field in self:
nfk = same_db or not field.type.startswith('reference') nfk = same_db or not field.type.startswith('reference')
clones.append(field.clone( clones.append(
unique=False, type=field.type if nfk else 'bigint')) field.clone(unique=False, type=field.type if nfk else 'bigint')
)
archive_db.define_table( archive_db.define_table(
archive_name, archive_name,
Field(current_record, field_type, label=current_record_label), Field(current_record, field_type, label=current_record_label),
@@ -8809,7 +8818,7 @@ class Table(object):
self._referenced_by = [] self._referenced_by = []
self._references = [] self._references = []
for field in self: for field in self:
fieldname = field.name #fieldname = field.name ##FIXME not used ?
field_type = field.type field_type = field.type
if isinstance(field_type, str) and field_type[:10] == 'reference ': if isinstance(field_type, str) and field_type[:10] == 'reference ':
ref = field_type[10:].strip() ref = field_type[10:].strip()
@@ -8829,8 +8838,9 @@ class Table(object):
'keyed tables can only reference other keyed tables (for now)') 'keyed tables can only reference other keyed tables (for now)')
if rfieldname not in rtable.fields: if rfieldname not in rtable.fields:
raise SyntaxError( raise SyntaxError(
"invalid field '%s' for referenced table '%s' in table '%s'" \ "invalid field '%s' for referenced table '%s'"
% (rfieldname, rtablename, self._tablename)) " in table '%s'" % (rfieldname, rtablename, self._tablename)
)
rfield = rtable[rfieldname] rfield = rtable[rfieldname]
else: else:
rfield = rtable._id rfield = rtable._id
@@ -8844,7 +8854,6 @@ class Table(object):
for referee in referees: for referee in referees:
self._referenced_by.append(referee) self._referenced_by.append(referee)
def _filter_fields(self, record, id=False): def _filter_fields(self, record, id=False):
return dict([(k, v) for (k, v) in record.iteritems() if k return dict([(k, v) for (k, v) in record.iteritems() if k
in self.fields and (self[k].type!='id' or id)]) in self.fields and (self[k].type!='id' or id)])
@@ -8860,8 +8869,9 @@ class Table(object):
query = (self[k] == v) query = (self[k] == v)
else: else:
raise SyntaxError( raise SyntaxError(
'Field %s is not part of the primary key of %s' % \ 'Field %s is not part of the primary key of %s' %
(k,self._tablename)) (k,self._tablename)
)
return query return query
def __getitem__(self, key): def __getitem__(self, key):
@@ -8878,10 +8888,12 @@ class Table(object):
def __call__(self, key=DEFAULT, **kwargs): def __call__(self, key=DEFAULT, **kwargs):
for_update = kwargs.get('_for_update', False) for_update = kwargs.get('_for_update', False)
if '_for_update' in kwargs: del kwargs['_for_update'] if '_for_update' in kwargs:
del kwargs['_for_update']
orderby = kwargs.get('_orderby', None) orderby = kwargs.get('_orderby', None)
if '_orderby' in kwargs: del kwargs['_orderby'] if '_orderby' in kwargs:
del kwargs['_orderby']
if not key is DEFAULT: if not key is DEFAULT:
if isinstance(key, Query): if isinstance(key, Query):
@@ -8915,7 +8927,7 @@ class Table(object):
self._db(query).update(**self._filter_fields(value)) self._db(query).update(**self._filter_fields(value))
else: else:
raise SyntaxError( raise SyntaxError(
'key must have all fields from primary key: %s'%\ 'key must have all fields from primary key: %s'%
(self._primarykey)) (self._primarykey))
elif str(key).isdigit(): elif str(key).isdigit():
if key == 0: if key == 0:
@@ -8960,7 +8972,6 @@ class Table(object):
def iteritems(self): def iteritems(self):
return self.__dict__.iteritems() return self.__dict__.iteritems()
def __repr__(self): def __repr__(self):
return '<Table %s (%s)>' % (self._tablename, ','.join(self.fields())) return '<Table %s (%s)>' % (self._tablename, ','.join(self.fields()))
@@ -9295,11 +9306,15 @@ class Table(object):
id_map_self[long(line[cid])] = new_id id_map_self[long(line[cid])] = new_id
def as_dict(self, flat=False, sanitize=True): def as_dict(self, flat=False, sanitize=True):
table_as_dict = dict(tablename=str(self), fields=[], table_as_dict = dict(
tablename=str(self),
fields=[],
sequence_name=self._sequence_name, sequence_name=self._sequence_name,
trigger_name=self._trigger_name, trigger_name=self._trigger_name,
common_filter=self._common_filter, format=self._format, common_filter=self._common_filter,
singular=self._singular, plural=self._plural) format=self._format,
singular=self._singular,
plural=self._plural)
for field in self: for field in self:
if (field.readable or field.writable) or (not sanitize): if (field.readable or field.writable) or (not sanitize):
@@ -9331,10 +9346,11 @@ class Table(object):
def on(self, query): def on(self, query):
return Expression(self._db, self._db._adapter.ON, self, query) return Expression(self._db, self._db._adapter.ON, self, query)
def archive_record(qset, fs, archive_table, current_record): def archive_record(qset, fs, archive_table, current_record):
tablenames = qset.db._adapter.tables(qset.query) tablenames = qset.db._adapter.tables(qset.query)
if len(tablenames)!=1: raise RuntimeError("cannot update join") if len(tablenames) != 1:
table = qset.db[tablenames[0]] raise RuntimeError("cannot update join")
for row in qset.select(): for row in qset.select():
fields = archive_table._filter_fields(row) fields = archive_table._filter_fields(row)
fields[current_record] = row.id fields[current_record] = row.id
@@ -9342,7 +9358,6 @@ def archive_record(qset,fs,archive_table,current_record):
return False return False
class Expression(object): class Expression(object):
def __init__( def __init__(
@@ -9724,15 +9739,18 @@ class FieldVirtual(object):
def __str__(self): def __str__(self):
return '%s.%s' % (self.tablename, self.name) return '%s.%s' % (self.tablename, self.name)
class FieldMethod(object): class FieldMethod(object):
def __init__(self, name, f=None, handler=None): def __init__(self, name, f=None, handler=None):
# for backward compatibility # for backward compatibility
(self.name, self.f) = (name, f) if f else ('unknown', name) (self.name, self.f) = (name, f) if f else ('unknown', name)
self.handler = handler self.handler = handler
def list_represent(x,r=None): def list_represent(x,r=None):
return ', '.join(str(y) for y in x or []) return ', '.join(str(y) for y in x or [])
class Field(Expression): class Field(Expression):
Virtual = FieldVirtual Virtual = FieldVirtual
@@ -9815,8 +9833,10 @@ class Field(Expression):
raise SyntaxError('Field: invalid unicode field name') raise SyntaxError('Field: invalid unicode field name')
self.name = fieldname = cleanup(fieldname) self.name = fieldname = cleanup(fieldname)
if not isinstance(fieldname, str) or hasattr(Table, fieldname) or \ if not isinstance(fieldname, str) or hasattr(Table, fieldname) or \
fieldname[0] == '_' or REGEX_PYTHON_KEYWORDS.match(fieldname): fieldname[0] == '_' or '.' in fieldname or \
raise SyntaxError('Field: invalid field name: %s' % fieldname) REGEX_PYTHON_KEYWORDS.match(fieldname):
raise SyntaxError('Field: invalid field name: %s, '
'use rname for "funny" names' % fieldname)
if not isinstance(type, (Table, Field)): if not isinstance(type, (Table, Field)):
self.type = type self.type = type
@@ -9840,8 +9860,8 @@ class Field(Expression):
self.update = update self.update = update
self.authorize = authorize self.authorize = authorize
self.autodelete = autodelete self.autodelete = autodelete
self.represent = list_represent if \ self.represent = (list_represent if represent is None and
represent==None and type in ('list:integer','list:string') else represent type in ('list:integer', 'list:string') else represent)
self.compute = compute self.compute = compute
self.isattachment = True self.isattachment = True
self.custom_store = custom_store self.custom_store = custom_store
@@ -9851,8 +9871,9 @@ class Field(Expression):
self.filter_in = filter_in self.filter_in = filter_in
self.filter_out = filter_out self.filter_out = filter_out
self.custom_qualifier = custom_qualifier self.custom_qualifier = custom_qualifier
self.label = label if label!=None else fieldname.replace('_',' ').title() self.label = (label if label is not None else
self.requires = requires if requires!=None else [] fieldname.replace('_', ' ').title())
self.requires = requires if requires is not None else []
self.map_none = map_none self.map_none = map_none
self._rname = rname self._rname = rname
@@ -9875,8 +9896,7 @@ class Field(Expression):
file = file.file file = file.file
elif not filename: elif not filename:
filename = file.name filename = file.name
filename = os.path.basename(filename.replace('/', os.sep)\ filename = os.path.basename(filename.replace('/', os.sep).replace('\\', os.sep))
.replace('\\', os.sep))
m = REGEX_STORE_PATTERN.search(filename) m = REGEX_STORE_PATTERN.search(filename)
extension = m and m.group('e') or 'txt' extension = m and m.group('e') or 'txt'
uuid_key = web2py_uuid().replace('-', '')[-16:] uuid_key = web2py_uuid().replace('-', '')[-16:]
@@ -9890,7 +9910,7 @@ class Field(Expression):
keys = {self_uploadfield.name: newfilename, keys = {self_uploadfield.name: newfilename,
blob_uploadfield_name: file.read()} blob_uploadfield_name: file.read()}
self_uploadfield.table.insert(**keys) self_uploadfield.table.insert(**keys)
elif self_uploadfield == True: elif self_uploadfield is True:
if path: if path:
pass pass
elif self.uploadfolder: elif self.uploadfolder:
@@ -9903,8 +9923,9 @@ class Field(Expression):
if self.uploadseparate: if self.uploadseparate:
if self.uploadfs: if self.uploadfs:
raise RuntimeError("not supported") raise RuntimeError("not supported")
path = pjoin(path,"%s.%s" %(self._tablename, self.name), path = pjoin(path, "%s.%s" % (
uuid_key[:2]) self._tablename, self.name), uuid_key[:2]
)
if not exists(path): if not exists(path):
os.makedirs(path) os.makedirs(path)
pathfilename = pjoin(path, newfilename) pathfilename = pjoin(path, newfilename)
@@ -9916,7 +9937,8 @@ class Field(Expression):
shutil.copyfileobj(file, dest_file) shutil.copyfileobj(file, dest_file)
except IOError: except IOError:
raise IOError( raise IOError(
'Unable to store file "%s" because invalid permissions, readonly file system, or filename too long' % pathfilename) 'Unable to store file "%s" because invalid permissions, '
'readonly file system, or filename too long' % pathfilename)
dest_file.close() dest_file.close()
return newfilename return newfilename
@@ -9988,7 +10010,6 @@ class Field(Expression):
path = pjoin(path, "%s.%s" % (t, f), u[:2]) path = pjoin(path, "%s.%s" % (t, f), u[:2])
return dict(path=path, filename=filename) return dict(path=path, filename=filename)
def formatter(self, value): def formatter(self, value):
requires = self.requires requires = self.requires
if value is None or not requires: if value is None or not requires:
@@ -10021,7 +10042,8 @@ class Field(Expression):
return Expression(self.db, self.db._adapter.COUNT, self, distinct, 'integer') return Expression(self.db, self.db._adapter.COUNT, self, distinct, 'integer')
def as_dict(self, flat=False, sanitize=True): def as_dict(self, flat=False, sanitize=True):
attrs = ('name', 'authorize', 'represent', 'ondelete', attrs = (
'name', 'authorize', 'represent', 'ondelete',
'custom_store', 'autodelete', 'custom_retrieve', 'custom_store', 'autodelete', 'custom_retrieve',
'filter_out', 'uploadseparate', 'widget', 'uploadfs', 'filter_out', 'uploadseparate', 'widget', 'uploadfs',
'update', 'custom_delete', 'uploadfield', 'uploadfolder', 'update', 'custom_delete', 'uploadfield', 'uploadfolder',
@@ -10034,8 +10056,7 @@ class Field(Expression):
def flatten(obj): def flatten(obj):
if isinstance(obj, dict): if isinstance(obj, dict):
return dict((flatten(k), flatten(v)) for k, v in return dict((flatten(k), flatten(v)) for k, v in obj.items())
obj.items())
elif isinstance(obj, (tuple, list, set)): elif isinstance(obj, (tuple, list, set)):
return [flatten(v) for v in obj] return [flatten(v) for v in obj]
elif isinstance(obj, serializable): elif isinstance(obj, serializable):
@@ -10177,6 +10198,7 @@ class Query(object):
SERIALIZABLE_TYPES = (tuple, dict, set, list, int, long, float, SERIALIZABLE_TYPES = (tuple, dict, set, list, int, long, float,
basestring, type(None), bool) basestring, type(None), bool)
def loop(d): def loop(d):
newd = dict() newd = dict()
for k, v in d.items(): for k, v in d.items():
@@ -10210,7 +10232,6 @@ class Query(object):
return loop(self.__dict__) return loop(self.__dict__)
else: return self.__dict__ else: return self.__dict__
def as_xml(self, sanitize=True): def as_xml(self, sanitize=True):
if have_serializers: if have_serializers:
xml = serializers.xml xml = serializers.xml
@@ -10227,6 +10248,7 @@ class Query(object):
d = self.as_dict(flat=True, sanitize=sanitize) d = self.as_dict(flat=True, sanitize=sanitize)
return json(d) return json(d)
def xorify(orderby): def xorify(orderby):
if not orderby: if not orderby:
return None return None
@@ -10235,10 +10257,12 @@ def xorify(orderby):
orderby2 = orderby2 | item orderby2 = orderby2 | item
return orderby2 return orderby2
def use_common_filters(query): def use_common_filters(query):
return (query and hasattr(query,'ignore_common_filters') and \ return (query and hasattr(query,'ignore_common_filters') and \
not query.ignore_common_filters) not query.ignore_common_filters)
class Set(object): class Set(object):
""" """
@@ -10850,7 +10874,6 @@ class Rows(object):
"represent" attributes will be transformed). "represent" attributes will be transformed).
""" """
if i is None: if i is None:
return (self.render(i, fields=fields) for i in range(len(self))) return (self.render(i, fields=fields) for i in range(len(self)))
import sqlhtml import sqlhtml
@@ -10889,7 +10912,6 @@ class Rows(object):
self.compact = compact self.compact = compact
return items return items
def as_dict(self, def as_dict(self,
key='id', key='id',
compact=True, compact=True,
@@ -10929,7 +10951,6 @@ class Rows(object):
else: else:
return dict([(key(r),r) for r in rows]) return dict([(key(r),r) for r in rows])
def as_trees(self, parent_name='parent_id', children_name='children'): def as_trees(self, parent_name='parent_id', children_name='children'):
roots = [] roots = []
drows = {} drows = {}
@@ -10964,6 +10985,7 @@ class Rows(object):
represent = kwargs.get('represent', False) represent = kwargs.get('represent', False)
writer = csv.writer(ofile, delimiter=delimiter, writer = csv.writer(ofile, delimiter=delimiter,
quotechar=quotechar, quoting=quoting) quotechar=quotechar, quoting=quoting)
def unquote_colnames(colnames): def unquote_colnames(colnames):
unq_colnames = [] unq_colnames = []
for col in colnames: for col in colnames:
+37 -41
View File
@@ -79,9 +79,6 @@ def tearDownModule():
class TestFields(unittest.TestCase): class TestFields(unittest.TestCase):
def testFieldName(self): def testFieldName(self):
return
# Any table name is supported as long as underlying db does. The following code is ignored.
# Check that Fields cannot start with underscores # Check that Fields cannot start with underscores
self.assertRaises(SyntaxError, Field, '_abc', 'string') self.assertRaises(SyntaxError, Field, '_abc', 'string')
@@ -201,6 +198,25 @@ class TestFields(unittest.TestCase):
db.tt.drop() db.tt.drop()
class TestTables(unittest.TestCase):
def testTableNames(self):
# Check that Tables cannot start with underscores
self.assertRaises(SyntaxError, Table, None, '_abc')
# Check that Tables cannot contain punctuation other than underscores
self.assertRaises(SyntaxError, Table, None, 'a.bc')
# Check that Tables cannot be a name of a method or property of DAL
for x in ['define_table', 'tables', 'as_dict']:
self.assertRaises(SyntaxError, Table, None, x)
# Check that Table allows underscores in the body of a field name.
self.assert_(Table(None, 'a_bc'),
"Table isn't allowing underscores in tablename. It should.")
class TestAll(unittest.TestCase): class TestAll(unittest.TestCase):
def setUp(self): def setUp(self):
@@ -1396,45 +1412,7 @@ class TestRNameFields(unittest.TestCase):
self.assertEqual(len(db.person._referenced_by),0) self.assertEqual(len(db.person._referenced_by),0)
db.person.drop() db.person.drop()
class TestQuoting(unittest.TestCase): class TestQuoting(unittest.TestCase):
# tests for complex table names
def testRun(self):
return
db = DAL(DEFAULT_URI, check_reserved=['all'])
t0 = db.define_table('A.table.with.dots and spaces',
Field('f', 'string'))
t1 = db.define_table('A.table',
Field('f_other', t0),
Field('words', 'text'))
blather = 'blah blah and so'
t0[0] = {'f': 'content'}
t1[0] = {'f_other': int(t0[1]['id']),
'words': blather}
r = db(t1['f_other']==t0.id).select()
self.assertEqual(r[0][db['A.table']].words, blather)
db.define_table('t0', Field('f0'))
db.define_table('t1', Field('f1'), Field('t0', db['t0']))
db.t0[0]=dict(f0=3)
db.t1[0]=dict(f1=3, t0=1)
rows=db(db.t0.id==db.t1.t0).select()
self.assertEqual(rows[0].t1.t0, rows[0].t0.id)
if DEFAULT_URI.startswith('mssql'):
#there's no drop cascade in mssql
t1.drop()
t0.drop()
else:
t0.drop('cascade')
t1.drop()
db.t1.drop()
db.t0.drop()
# tests for case sensitivity # tests for case sensitivity
def testCase(self): def testCase(self):
@@ -1528,6 +1506,24 @@ class TestQuoting(unittest.TestCase):
t3.drop() t3.drop()
t4.drop() t4.drop()
class TestTableAndFieldCase(unittest.TestCase):
"""
at the Python level we should not allow db.C and db.c because of .table conflicts on windows
but it should be possible to map two different names into distinct tables "c" and "C" at the Python level
By default Python models names should be mapped into lower case table names and assume case insensitivity.
"""
def testme(self):
return
class TestQuotesByDefault(unittest.TestCase):
"""
all default tables names should be quoted unless an explicit mapping has been given for a table.
"""
def testme(self):
return
if __name__ == '__main__': if __name__ == '__main__':
unittest.main() unittest.main()
tearDownModule() tearDownModule()