Merge pull request #317 from michele-comitini/entity_quoting

quoting of tablenames and field names
This commit is contained in:
mdipierro
2013-12-04 08:34:52 -08:00
3 changed files with 389 additions and 156 deletions
+274 -152
View File
@@ -252,7 +252,8 @@ THREAD_LOCAL = threading.local()
REGEX_TYPE = re.compile('^([\w\_\:]+)') REGEX_TYPE = re.compile('^([\w\_\:]+)')
REGEX_DBNAME = re.compile('^(\w+)(\:\w+)*') REGEX_DBNAME = re.compile('^(\w+)(\:\w+)*')
REGEX_W = re.compile('^\w+$') REGEX_W = re.compile('^\w+$')
REGEX_TABLE_DOT_FIELD = re.compile('^(\w+)\.(\w+)$') REGEX_TABLE_DOT_FIELD = re.compile('^(\w+)\.([^.]+)$')
REGEX_NO_GREEDY_ENTITY_NAME = r'(.+?)'
REGEX_UPLOAD_PATTERN = re.compile('(?P<table>[\w\-]+)\.(?P<field>[\w\-]+)\.(?P<uuidkey>[\w\-]+)(\.(?P<name>\w+))?\.\w+$') REGEX_UPLOAD_PATTERN = re.compile('(?P<table>[\w\-]+)\.(?P<field>[\w\-]+)\.(?P<uuidkey>[\w\-]+)(\.(?P<name>\w+))?\.\w+$')
REGEX_CLEANUP_FN = re.compile('[\'"\s;]+') REGEX_CLEANUP_FN = re.compile('[\'"\s;]+')
REGEX_UNPACK = re.compile('(?<!\|)\|(?!\|)') REGEX_UNPACK = re.compile('(?<!\|)\|(?!\|)')
@@ -670,12 +671,27 @@ class ConnectionPool(object):
break break
self.after_connection_hook() self.after_connection_hook()
###################################################################################
# metaclass to prepare adapter classes static values
###################################################################################
class AdapterMeta(type):
def __new__(cls, clsname, bases, dct):
classobj = super(AdapterMeta, cls).__new__(cls, clsname, bases, dct)
classobj.REGEX_TABLE_DOT_FIELD = re.compile(r'^' + \
classobj.QUOTE_TEMPLATE % REGEX_NO_GREEDY_ENTITY_NAME + \
r'\.' + \
classobj.QUOTE_TEMPLATE % REGEX_NO_GREEDY_ENTITY_NAME + \
r'$')
return classobj
################################################################################### ###################################################################################
# this is a generic adapter that does nothing; all others are derived from this one # this is a generic adapter that does nothing; all others are derived from this one
################################################################################### ###################################################################################
class BaseAdapter(ConnectionPool): class BaseAdapter(ConnectionPool):
__metaclass__ = AdapterMeta
native_json = False native_json = False
driver = None driver = None
driver_name = None driver_name = None
@@ -693,6 +709,7 @@ class BaseAdapter(ConnectionPool):
T_SEP = ' ' T_SEP = ' '
QUOTE_TEMPLATE = '"%s"' QUOTE_TEMPLATE = '"%s"'
types = { types = {
'boolean': 'CHAR(1)', 'boolean': 'CHAR(1)',
'string': 'CHAR(%(length)s)', 'string': 'CHAR(%(length)s)',
@@ -717,6 +734,7 @@ class BaseAdapter(ConnectionPool):
# the two below are only used when DAL(...bigint_id=True) and replace 'id','reference' # the two below are only used when DAL(...bigint_id=True) and replace 'id','reference'
'big-id': 'BIGINT PRIMARY KEY AUTOINCREMENT', 'big-id': 'BIGINT PRIMARY KEY AUTOINCREMENT',
'big-reference': 'BIGINT REFERENCES %(foreign_key)s ON DELETE %(on_delete_action)s', 'big-reference': 'BIGINT REFERENCES %(foreign_key)s ON DELETE %(on_delete_action)s',
'reference FK': ', CONSTRAINT "FK_%(constraint_name)s" FOREIGN KEY (%(field_name)s) REFERENCES %(foreign_key)s ON DELETE %(on_delete_action)s',
} }
def isOperationalError(self,exception): def isOperationalError(self,exception):
@@ -834,8 +852,9 @@ class BaseAdapter(ConnectionPool):
self.connection = Dummy() self.connection = Dummy()
self.cursor = Dummy() self.cursor = Dummy()
def sequence_name(self,tablename): def sequence_name(self,tablename):
return '%s_sequence' % tablename return self.QUOTE_TEMPLATE % ('%s_sequence' % tablename)
def trigger_name(self,tablename): def trigger_name(self,tablename):
return '%s_sequence' % tablename return '%s_sequence' % tablename
@@ -868,61 +887,72 @@ class BaseAdapter(ConnectionPool):
if referenced == '.': if referenced == '.':
referenced = tablename referenced = tablename
constraint_name = self.constraint_name(tablename, field_name) constraint_name = self.constraint_name(tablename, field_name)
if not '.' in referenced \ # if not '.' in referenced \
and referenced != tablename \ # and referenced != tablename \
and hasattr(table,'_primarykey'): # and hasattr(table,'_primarykey'):
ftype = types['integer'] # ftype = types['integer']
else: #else:
if hasattr(table,'_primarykey'): try:
rtable = db[referenced]
rfield = rtable._id
rfieldname = rfield.name
rtablename = referenced
except (KeyError, ValueError, AttributeError), e:
LOGGER.debug('Error: %s' % e)
try:
rtablename,rfieldname = referenced.split('.') rtablename,rfieldname = referenced.split('.')
rtable = db[rtablename] rtable = db[rtablename]
rfield = rtable[rfieldname] rfield = rtable[rfieldname]
# must be PK reference or unique except Exception, e:
if rfieldname in rtable._primarykey or \ LOGGER.debug('Error: %s' %e)
rfield.unique: raise KeyError('Cannot resolve reference %s in %s definition' % (referenced, table._tablename))
ftype = types[rfield.type[:9]] % \
dict(length=rfield.length) # must be PK reference or unique
# multicolumn primary key reference? if getattr(rtable, '_primarykey', None) and rfieldname in rtable._primarykey or \
if not rfield.unique and len(rtable._primarykey)>1: rfield.unique:
# then it has to be a table level FK ftype = types[rfield.type[:9]] % \
if rtablename not in TFK: dict(length=rfield.length)
TFK[rtablename] = {} # multicolumn primary key reference?
TFK[rtablename][rfieldname] = field_name if not rfield.unique and len(rtable._primarykey)>1:
else: # then it has to be a table level FK
ftype = ftype + \ if rtablename not in TFK:
types['reference FK'] % dict( TFK[rtablename] = {}
constraint_name = constraint_name, # should be quoted TFK[rtablename][rfieldname] = field_name
foreign_key = '%s (%s)' % (rtablename,
rfieldname),
table_name = tablename,
field_name = field._rname or field.name,
on_delete_action=field.ondelete)
else: else:
# make a guess here for circular references ftype = ftype + \
if referenced in db: types['reference FK'] % dict(
id_fieldname = db[referenced]._id.name constraint_name = constraint_name, # should be quoted
elif referenced == tablename: foreign_key = rtable.sqlsafe + ' (' + rfield.sqlsafe_name + ')',
id_fieldname = table._id.name table_name = table.sqlsafe,
else: #make a guess field_name = field.sqlsafe_name,
id_fieldname = 'id' on_delete_action=field.ondelete)
#gotcha: the referenced table must be defined before else:
#the referencing one to be able to create the table # make a guess here for circular references
#Also if it's not recommended, we can still support if referenced in db:
#references to tablenames without rname to make id_fieldname = db[referenced]._id.sqlsafe_name
#migrations and model relationship work also if tables elif referenced == tablename:
#are not defined in order id_fieldname = table._id.sqlsafe_name
real_referenced = ( else: #make a guess
(db[referenced]._rname or db[referenced]) id_fieldname = self.QUOTE_TEMPLATE % 'id'
if referenced == tablename or referenced in db #gotcha: the referenced table must be defined before
else referenced) #the referencing one to be able to create the table
#Also if it's not recommended, we can still support
ftype = types[field_type[:9]] % dict( #references to tablenames without rname to make
index_name = field_name+'__idx', #migrations and model relationship work also if tables
field_name = field._rname or field.name, #are not defined in order
constraint_name = constraint_name, if referenced == tablename:
foreign_key = '%s (%s)' % (real_referenced, real_referenced = db[referenced].sqlsafe
id_fieldname), else:
on_delete_action=field.ondelete) real_referenced = (referenced in db
and db[referenced].sqlsafe
or referenced)
rfield = db[referenced]._id
ftype = types[field_type[:9]] % dict(
index_name = self.QUOTE_TEMPLATE % (field_name+'__idx'),
field_name = field.sqlsafe_name,
constraint_name = self.QUOTE_TEMPLATE % constraint_name,
foreign_key = '%s (%s)' % (real_referenced, rfield.sqlsafe_name),
on_delete_action=field.ondelete)
elif field_type.startswith('list:reference'): elif field_type.startswith('list:reference'):
ftype = types[field_type[:14]] ftype = types[field_type[:14]]
elif field_type.startswith('decimal'): elif field_type.startswith('decimal'):
@@ -995,39 +1025,37 @@ class BaseAdapter(ConnectionPool):
# geometry fields are added after the table has been created, not now # geometry fields are added after the table has been created, not now
if not (self.dbengine == 'postgres' and \ if not (self.dbengine == 'postgres' and \
field_type.startswith('geom')): field_type.startswith('geom')):
#fetch the rname if it's there fields.append('%s %s' % (field.sqlsafe_name, ftype))
field_rname = "%s" % (field._rname or field_name)
fields.append('%s %s' % (field_rname, ftype))
other = ';' other = ';'
# backend-specific extensions to fields # backend-specific extensions to fields
if self.dbengine == 'mysql': if self.dbengine == 'mysql':
if not hasattr(table, "_primarykey"): if not hasattr(table, "_primarykey"):
fields.append('PRIMARY KEY(%s)' % (table._id.name or table._id._rname)) fields.append('PRIMARY KEY (%s)' % (self.QUOTE_TEMPLATE % table._id.name))
other = ' ENGINE=InnoDB CHARACTER SET utf8;' other = ' ENGINE=InnoDB CHARACTER SET utf8;'
fields = ',\n '.join(fields) fields = ',\n '.join(fields)
for rtablename in TFK: for rtablename in TFK:
rfields = TFK[rtablename] rfields = TFK[rtablename]
pkeys = db[rtablename]._primarykey pkeys = [self.QUOTE_TEMPLATE % pk for pk in db[rtablename]._primarykey]
fkeys = [ rfields[k] for k in pkeys ] fkeys = [self.QUOTE_TEMPLATE % rfields[k].name for k in pkeys ]
fields = fields + ',\n ' + \ fields = fields + ',\n ' + \
types['reference TFK'] % dict( types['reference TFK'] % dict(
table_name = tablename, table_name = table.sqlsafe,
field_name=', '.join(fkeys), field_name=', '.join(fkeys),
foreign_table = rtablename, foreign_table = table.sqlsafe,
foreign_key = ', '.join(pkeys), foreign_key = ', '.join(pkeys),
on_delete_action = field.ondelete) on_delete_action = field.ondelete)
#if there's a _rname, let's use that instead
table_rname = table._rname or tablename table_rname = table.sqlsafe
if getattr(table,'_primarykey',None): if getattr(table,'_primarykey',None):
query = "CREATE TABLE %s(\n %s,\n %s) %s" % \ query = "CREATE TABLE %s(\n %s,\n %s) %s" % \
(table_rname, fields, (table.sqlsafe, fields,
self.PRIMARY_KEY(', '.join(table._primarykey)),other) self.PRIMARY_KEY(', '.join([self.QUOTE_TEMPLATE % pk for pk in table._primarykey])),other)
else: else:
query = "CREATE TABLE %s(\n %s\n)%s" % \ query = "CREATE TABLE %s(\n %s\n)%s" % \
(table_rname, fields, other) (table.sqlsafe, fields, other)
if self.uri.startswith('sqlite:///') \ if self.uri.startswith('sqlite:///') \
or self.uri.startswith('spatialite:///'): or self.uri.startswith('spatialite:///'):
@@ -1105,6 +1133,7 @@ class BaseAdapter(ConnectionPool):
k,v=item k,v=item
if not isinstance(v,dict): if not isinstance(v,dict):
v=dict(type='unknown',sql=v) v=dict(type='unknown',sql=v)
if self.ignore_field_case is not True: return k, v
return k.lower(),v return k.lower(),v
# make sure all field names are lower case to avoid # make sure all field names are lower case to avoid
# migrations because of case cahnge # migrations because of case cahnge
@@ -1132,7 +1161,7 @@ class BaseAdapter(ConnectionPool):
query = [ sql_fields[key]['sql'] ] query = [ sql_fields[key]['sql'] ]
else: else:
query = ['ALTER TABLE %s ADD %s %s;' % \ query = ['ALTER TABLE %s ADD %s %s;' % \
(tablename, key, (table.sqlsafe, key,
sql_fields_aux[key]['sql'].replace(', ', new_add))] sql_fields_aux[key]['sql'].replace(', ', new_add))]
metadata_change = True metadata_change = True
elif self.dbengine in ('sqlite', 'spatialite'): elif self.dbengine in ('sqlite', 'spatialite'):
@@ -1150,10 +1179,11 @@ class BaseAdapter(ConnectionPool):
"'%(table)s', '%(field)s');" % "'%(table)s', '%(field)s');" %
dict(schema=schema, table=tablename, field=key,) ] dict(schema=schema, table=tablename, field=key,) ]
elif self.dbengine in ('firebird',): elif self.dbengine in ('firebird',):
query = ['ALTER TABLE %s DROP %s;' % (tablename, key)] query = ['ALTER TABLE %s DROP %s;' %
(self.QUOTE_TEMPLATE % tablename, self.QUOTE_TEMPLATE % key)]
else: else:
query = ['ALTER TABLE %s DROP COLUMN %s;' % query = ['ALTER TABLE %s DROP COLUMN %s;' %
(tablename, key)] (self.QUOTE_TEMPLATE % tablename, self.QUOTE_TEMPLATE % key)]
metadata_change = True metadata_change = True
elif sql_fields[key]['sql'] != sql_fields_old[key]['sql'] \ elif sql_fields[key]['sql'] != sql_fields_old[key]['sql'] \
and not (key in table.fields and and not (key in table.fields and
@@ -1169,12 +1199,15 @@ class BaseAdapter(ConnectionPool):
else: else:
drop_expr = 'ALTER TABLE %s DROP COLUMN %s;' drop_expr = 'ALTER TABLE %s DROP COLUMN %s;'
key_tmp = key + '__tmp' key_tmp = key + '__tmp'
query = ['ALTER TABLE %s ADD %s %s;' % (t, key_tmp, tt), query = ['ALTER TABLE %s ADD %s %s;' % (self.QUOTE_TEMPLATE % t, self.QUOTE_TEMPLATE % key_tmp, tt),
'UPDATE %s SET %s=%s;' % (t, key_tmp, key), 'UPDATE %s SET %s=%s;' %
drop_expr % (t, key), (self.QUOTE_TEMPLATE % t, self.QUOTE_TEMPLATE % key_tmp, self.QUOTE_TEMPLATE % key),
'ALTER TABLE %s ADD %s %s;' % (t, key, tt), drop_expr % (self.QUOTE_TEMPLATE % t, self.QUOTE_TEMPLATE % key),
'UPDATE %s SET %s=%s;' % (t, key, key_tmp), 'ALTER TABLE %s ADD %s %s;' %
drop_expr % (t, key_tmp)] (self.QUOTE_TEMPLATE % t, self.QUOTE_TEMPLATE % key, tt),
'UPDATE %s SET %s=%s;' %
(self.QUOTE_TEMPLATE % t, self.QUOTE_TEMPLATE % key, self.QUOTE_TEMPLATE % key_tmp),
drop_expr % (self.QUOTE_TEMPLATE % t, self.QUOTE_TEMPLATE % key_tmp)]
metadata_change = True metadata_change = True
elif sql_fields[key]['type'] != sql_fields_old[key]['type']: elif sql_fields[key]['type'] != sql_fields_old[key]['type']:
sql_fields_current[key] = sql_fields[key] sql_fields_current[key] = sql_fields[key]
@@ -1269,8 +1302,7 @@ class BaseAdapter(ConnectionPool):
return 'PRIMARY KEY(%s)' % key return 'PRIMARY KEY(%s)' % key
def _drop(self, table, mode): def _drop(self, table, mode):
table_rname = table._rname or table return ['DROP TABLE %s;' % table.sqlsafe]
return ['DROP TABLE %s;' % table_rname]
def drop(self, table, mode=''): def drop(self, table, mode=''):
db = table._db db = table._db
@@ -1288,17 +1320,16 @@ class BaseAdapter(ConnectionPool):
self.log('success!\n', table) self.log('success!\n', table)
def _insert(self, table, fields): def _insert(self, table, fields):
table_rname = table._rname or table table_rname = table.sqlsafe
if fields: if fields:
keys = ','.join(f._rname or f.name for f, v in fields) keys = ','.join(f.sqlsafe_name for f, v in fields)
values = ','.join(self.expand(v, f.type) for f, v in fields) values = ','.join(self.expand(v, f.type) for f, v in fields)
return 'INSERT INTO %s(%s) VALUES (%s);' % (table_rname, keys, values) return 'INSERT INTO %s(%s) VALUES (%s);' % (table_rname, keys, values)
else: else:
return self._insert_empty(table) return self._insert_empty(table)
def _insert_empty(self, table): def _insert_empty(self, table):
table_rname = table._rname or table return 'INSERT INTO %s DEFAULT VALUES;' % (table.sqlsafe)
return 'INSERT INTO %s DEFAULT VALUES;' % table_rname
def insert(self, table, fields): def insert(self, table, fields):
query = self._insert(table,fields) query = self._insert(table,fields)
@@ -1451,13 +1482,13 @@ class BaseAdapter(ConnectionPool):
self.expand(second, first.type)) self.expand(second, first.type))
def AS(self, first, second): def AS(self, first, second):
return '%s AS %s' % (self.expand(first), second) return '%s AS %s' % (self.expand(first), second)
def ON(self, first, second): def ON(self, first, second):
table_rname = first._ot and first or first._rname or first._tablename table_rname = self.table_alias(first)
if use_common_filters(second): if use_common_filters(second):
second = self.common_filter(second,[first._tablename]) second = self.common_filter(second,[first._tablename])
return '%s ON %s' % (self.expand(table_rname), self.expand(second)) return ('%s ON %s') % (self.expand(table_rname), self.expand(second))
def INVERT(self, first): def INVERT(self, first):
return '%s DESC' % self.expand(first) return '%s DESC' % self.expand(first)
@@ -1472,13 +1503,10 @@ class BaseAdapter(ConnectionPool):
if isinstance(expression, Field): if isinstance(expression, Field):
et = expression.table et = expression.table
if not colnames: if not colnames:
table_rname = et._ot and et._tablename or et._rname or et._tablename table_rname = et._ot and self.QUOTE_TEMPLATE % et._tablename or et._rname or self.QUOTE_TEMPLATE % et._tablename
out = '%s.%s' % (table_rname, expression._rname or (self.QUOTE_TEMPLATE % (expression.name)))
else: else:
table_rname = et._tablename out = '%s.%s' % (self.QUOTE_TEMPLATE % et._tablename, self.QUOTE_TEMPLATE % expression.name)
if not colnames:
out = '%s.%s' % (table_rname, expression._rname or expression.name)
else:
out = '%s.%s' % (table_rname, expression.name)
if field_type == 'string' and not expression.type in ( if field_type == 'string' and not expression.type in (
'string','text','json','password'): 'string','text','json','password'):
out = self.CAST(out, self.types['text']) out = self.CAST(out, self.types['text'])
@@ -1509,10 +1537,11 @@ class BaseAdapter(ConnectionPool):
else: else:
return str(expression) return str(expression)
def table_alias(self,name): def table_alias(self, tbl):
if not isinstance(name, Table): if not isinstance(tbl, Table):
name = self.db[name]._rname or self.db[name] tbl = self.db[tbl]
return str(name) return tbl.sqlsafe_alias
def alias(self, table, alias): def alias(self, table, alias):
""" """
@@ -1520,7 +1549,7 @@ class BaseAdapter(ConnectionPool):
with alias name. with alias name.
""" """
other = copy.copy(table) other = copy.copy(table)
other['_ot'] = other._ot or other._rname or other._tablename other['_ot'] = other._ot or other.sqlsafe
other['ALL'] = SQLALL(other) other['ALL'] = SQLALL(other)
other['_tablename'] = alias other['_tablename'] = alias
for fieldname in other.fields: for fieldname in other.fields:
@@ -1532,8 +1561,7 @@ class BaseAdapter(ConnectionPool):
return other return other
def _truncate(self, table, mode=''): def _truncate(self, table, mode=''):
tablename = table._rname or table._tablename return ['TRUNCATE TABLE %s %s;' % (table.sqlsafe, mode or '')]
return ['TRUNCATE TABLE %s %s;' % (tablename, mode or '')]
def truncate(self, table, mode= ' '): def truncate(self, table, mode= ' '):
# Prepare functions "write_to_logfile" and "close_logfile" # Prepare functions "write_to_logfile" and "close_logfile"
@@ -1553,10 +1581,10 @@ class BaseAdapter(ConnectionPool):
sql_w = ' WHERE ' + self.expand(query) sql_w = ' WHERE ' + self.expand(query)
else: else:
sql_w = '' sql_w = ''
sql_v = ','.join(['%s=%s' % (field._rname or field.name, sql_v = ','.join(['%s=%s' % (field.sqlsafe_name,
self.expand(value, field.type)) \ self.expand(value, field.type)) \
for (field, value) in fields]) for (field, value) in fields])
tablename = "%s" % (self.db[tablename]._rname or tablename) tablename = self.db[tablename].sqlsafe
return 'UPDATE %s SET %s%s;' % (tablename, sql_v, sql_w) return 'UPDATE %s SET %s%s;' % (tablename, sql_v, sql_w)
def update(self, tablename, query, fields): def update(self, tablename, query, fields):
@@ -1581,7 +1609,7 @@ class BaseAdapter(ConnectionPool):
sql_w = ' WHERE ' + self.expand(query) sql_w = ' WHERE ' + self.expand(query)
else: else:
sql_w = '' sql_w = ''
tablename = '%s' % (self.db[tablename]._rname or tablename) tablename = self.db[tablename].sqlsafe
return 'DELETE FROM %s%s;' % (tablename, sql_w) return 'DELETE FROM %s%s;' % (tablename, sql_w)
def delete(self, tablename, query): def delete(self, tablename, query):
@@ -1623,8 +1651,9 @@ class BaseAdapter(ConnectionPool):
if isinstance(item,SQLALL): if isinstance(item,SQLALL):
new_fields += item._table new_fields += item._table
elif isinstance(item,str): elif isinstance(item,str):
if REGEX_TABLE_DOT_FIELD.match(item): m = self.REGEX_TABLE_DOT_FIELD.match(item)
tablename,fieldname = item.split('.') if m:
tablename,fieldname = m.groups()
append(db[tablename][fieldname]) append(db[tablename][fieldname])
else: else:
append(Expression(db,lambda item=item:item)) append(Expression(db,lambda item=item:item))
@@ -1645,10 +1674,11 @@ class BaseAdapter(ConnectionPool):
tablenames = tables(query) tablenames = tables(query)
tablenames_for_common_filters = tablenames tablenames_for_common_filters = tablenames
for field in fields: for field in fields:
if isinstance(field, basestring) \ if isinstance(field, basestring):
and REGEX_TABLE_DOT_FIELD.match(field): m = self.REGEX_TABLE_DOT_FIELD.match(field)
tn,fn = field.split('.') if m:
field = self.db[tn][fn] tn,fn = m.groups()
field = self.db[tn][fn]
for tablename in tables(field): for tablename in tables(field):
if not tablename in tablenames: if not tablename in tablenames:
tablenames.append(tablename) tablenames.append(tablename)
@@ -1732,7 +1762,7 @@ class BaseAdapter(ConnectionPool):
tables_to_merge.keys()]) tables_to_merge.keys()])
if joint: if joint:
sql_t += ' %s %s' % (command, sql_t += ' %s %s' % (command,
','.join([self.table_alias(t) for t in joint])) ','.join([t for t in joint]))
for t in joinon: for t in joinon:
sql_t += ' %s %s' % (command, t) sql_t += ' %s %s' % (command, t)
elif inner_join and left: elif inner_join and left:
@@ -1747,7 +1777,7 @@ class BaseAdapter(ConnectionPool):
sql_t += ' %s %s' % (icommand, t) sql_t += ' %s %s' % (icommand, t)
if joint: if joint:
sql_t += ' %s %s' % (command, sql_t += ' %s %s' % (command,
','.join([self.table_alias(t) for t in joint])) ','.join([t for t in joint]))
for t in joinon: for t in joinon:
sql_t += ' %s %s' % (command, t) sql_t += ' %s %s' % (command, t)
else: else:
@@ -1767,9 +1797,9 @@ class BaseAdapter(ConnectionPool):
sql_o += ' ORDER BY %s' % self.expand(orderby) sql_o += ' ORDER BY %s' % self.expand(orderby)
if (limitby and not groupby and tablenames and orderby_on_limitby and not orderby): if (limitby and not groupby and tablenames and orderby_on_limitby and not orderby):
sql_o += ' ORDER BY %s' % ', '.join( sql_o += ' ORDER BY %s' % ', '.join(
['%s.%s'%(t,x) for t in tablenames for x in ( [self.db[t][x].sqlsafe for t in tablenames for x in (
hasattr(self.db[t],'_primarykey') and self.db[t]._primarykey hasattr(self.db[t],'_primarykey') and self.db[t]._primarykey
or [self.db[t]._id._rname or self.db[t]._id.name] or ['_id']
) )
] ]
) )
@@ -1897,8 +1927,10 @@ class BaseAdapter(ConnectionPool):
def create_sequence_and_triggers(self, query, table, **args): def create_sequence_and_triggers(self, query, table, **args):
self.execute(query) self.execute(query)
def log_execute(self, *a, **b): def log_execute(self, *a, **b):
if not self.connection: raise ValueError(a[0])
if not self.connection: return None if not self.connection: return None
command = a[0] command = a[0]
if hasattr(self,'filter_sql_command'): if hasattr(self,'filter_sql_command'):
@@ -1955,8 +1987,17 @@ class BaseAdapter(ConnectionPool):
if field_is_type('decimal'): if field_is_type('decimal'):
return str(obj) return str(obj)
elif field_is_type('reference'): # reference elif field_is_type('reference'): # reference
if fieldtype.find('.')>0: # check for tablename first
return repr(obj) referenced = fieldtype[9:].strip()
if referenced in self.db.tables:
return str(long(obj))
p = referenced.partition('.')
if p[2] != '':
try:
ftype = self.db[p[0]][p[2]].type
return self.represent(obj, ftype)
except (ValueError, KeyError):
return repr(obj)
elif isinstance(obj, (Row, Reference)): elif isinstance(obj, (Row, Reference)):
return str(obj['id']) return str(obj['id'])
return str(long(obj)) return str(long(obj))
@@ -2162,14 +2203,15 @@ class BaseAdapter(ConnectionPool):
new_rows = [] new_rows = []
tmps = [] tmps = []
for colname in colnames: for colname in colnames:
if not REGEX_TABLE_DOT_FIELD.match(colname): col_m = self.REGEX_TABLE_DOT_FIELD.match(colname)
if not col_m:
tmps.append(None) tmps.append(None)
else: else:
(tablename, _the_sep_, fieldname) = colname.partition('.') tablename, fieldname = col_m.groups()
table = db[tablename] table = db[tablename]
field = table[fieldname] field = table[fieldname]
ft = field.type ft = field.type
tmps.append((tablename,fieldname,table,field,ft)) tmps.append((tablename, fieldname, table, field, ft))
for (i,row) in enumerate(rows): for (i,row) in enumerate(rows):
new_row = Row() new_row = Row()
for (j,colname) in enumerate(colnames): for (j,colname) in enumerate(colnames):
@@ -2177,9 +2219,8 @@ class BaseAdapter(ConnectionPool):
tmp = tmps[j] tmp = tmps[j]
if tmp: if tmp:
(tablename,fieldname,table,field,ft) = tmp (tablename,fieldname,table,field,ft) = tmp
if tablename in new_row: colset = new_row.get(tablename, None)
colset = new_row[tablename] if colset is None:
else:
colset = new_row[tablename] = Row() colset = new_row[tablename] = Row()
if tablename not in virtualtables: if tablename not in virtualtables:
virtualtables.append(tablename) virtualtables.append(tablename)
@@ -2287,6 +2328,14 @@ class BaseAdapter(ConnectionPool):
return Expression(self.db,'CASE WHEN %s THEN %s ELSE %s END' % \ return Expression(self.db,'CASE WHEN %s THEN %s ELSE %s END' % \
(self.expand(query),represent(t),represent(f))) (self.expand(query),represent(t),represent(f)))
def sqlsafe_table(self, tablename, ot=None):
if ot is not None:
return ('%s AS ' + self.QUOTE_TEMPLATE) % (ot, tablename)
return self.QUOTE_TEMPLATE % tablename
def sqlsafe_field(self, fieldname):
return self.QUOTE_TEMPLATE % fieldname
################################################################################### ###################################################################################
# List of all the available adapters; they all extend BaseAdapter. # List of all the available adapters; they all extend BaseAdapter.
################################################################################### ###################################################################################
@@ -2564,6 +2613,7 @@ class MySQLAdapter(BaseAdapter):
'list:reference': 'LONGTEXT', 'list:reference': 'LONGTEXT',
'big-id': 'BIGINT AUTO_INCREMENT NOT NULL', 'big-id': 'BIGINT AUTO_INCREMENT NOT NULL',
'big-reference': 'BIGINT, INDEX %(index_name)s (%(field_name)s), FOREIGN KEY (%(field_name)s) REFERENCES %(foreign_key)s ON DELETE %(on_delete_action)s', 'big-reference': 'BIGINT, INDEX %(index_name)s (%(field_name)s), FOREIGN KEY (%(field_name)s) REFERENCES %(foreign_key)s ON DELETE %(on_delete_action)s',
'reference FK': ', CONSTRAINT `FK_%(constraint_name)s` FOREIGN KEY (%(field_name)s) REFERENCES %(foreign_key)s ON DELETE %(on_delete_action)s',
} }
QUOTE_TEMPLATE = "`%s`" QUOTE_TEMPLATE = "`%s`"
@@ -2590,12 +2640,12 @@ class MySQLAdapter(BaseAdapter):
def _drop(self,table,mode): def _drop(self,table,mode):
# breaks db integrity but without this mysql does not drop table # breaks db integrity but without this mysql does not drop table
table_rname = table._rname or table table_rname = table.sqlsafe
return ['SET FOREIGN_KEY_CHECKS=0;','DROP TABLE %s;' % table_rname, return ['SET FOREIGN_KEY_CHECKS=0;','DROP TABLE %s;' % table_rname,
'SET FOREIGN_KEY_CHECKS=1;'] 'SET FOREIGN_KEY_CHECKS=1;']
def _insert_empty(self, table): def _insert_empty(self, table):
return 'INSERT INTO %s VALUES (DEFAULT);' % table return 'INSERT INTO %s VALUES (DEFAULT);' % (table.sqlsafe)
def distributed_transaction_begin(self,key): def distributed_transaction_begin(self,key):
self.execute('XA START;') self.execute('XA START;')
@@ -2668,6 +2718,8 @@ class MySQLAdapter(BaseAdapter):
class PostgreSQLAdapter(BaseAdapter): class PostgreSQLAdapter(BaseAdapter):
drivers = ('psycopg2','pg8000') drivers = ('psycopg2','pg8000')
QUOTE_TEMPLATE = '"%s"'
support_distributed_transaction = True support_distributed_transaction = True
types = { types = {
'boolean': 'CHAR(1)', 'boolean': 'CHAR(1)',
@@ -2694,12 +2746,11 @@ class PostgreSQLAdapter(BaseAdapter):
'geography': 'GEOGRAPHY', 'geography': 'GEOGRAPHY',
'big-id': 'BIGSERIAL PRIMARY KEY', 'big-id': 'BIGSERIAL PRIMARY KEY',
'big-reference': 'BIGINT REFERENCES %(foreign_key)s ON DELETE %(on_delete_action)s', 'big-reference': 'BIGINT REFERENCES %(foreign_key)s ON DELETE %(on_delete_action)s',
'reference FK': ', CONSTRAINT FK_%(constraint_name)s FOREIGN KEY (%(field_name)s) REFERENCES %(foreign_key)s ON DELETE %(on_delete_action)s', 'reference FK': ', CONSTRAINT "FK_%(constraint_name)s" FOREIGN KEY (%(field_name)s) REFERENCES %(foreign_key)s ON DELETE %(on_delete_action)s',
'reference TFK': ' CONSTRAINT FK_%(foreign_table)s_PK FOREIGN KEY (%(field_name)s) REFERENCES %(foreign_table)s (%(foreign_key)s) ON DELETE %(on_delete_action)s', 'reference TFK': ' CONSTRAINT "FK_%(foreign_table)s_PK" FOREIGN KEY (%(field_name)s) REFERENCES %(foreign_table)s (%(foreign_key)s) ON DELETE %(on_delete_action)s',
} }
QUOTE_TEMPLATE = '"%s"'
def varquote(self,name): def varquote(self,name):
return varquote_aux(name,'"%s"') return varquote_aux(name,'"%s"')
@@ -2713,7 +2764,7 @@ class PostgreSQLAdapter(BaseAdapter):
return "'%s'" % str(obj).replace("'","''") return "'%s'" % str(obj).replace("'","''")
def sequence_name(self,table): def sequence_name(self,table):
return '%s_id_seq' % table return self.QUOTE_TEMPLATE % (table + '_id_seq')
def RANDOM(self): def RANDOM(self):
return 'RANDOM()' return 'RANDOM()'
@@ -2802,8 +2853,8 @@ class PostgreSQLAdapter(BaseAdapter):
self.execute("SET standard_conforming_strings=on;") self.execute("SET standard_conforming_strings=on;")
self.try_json() self.try_json()
def lastrowid(self,table): def lastrowid(self,table = None):
self.execute("""select currval('"%s"')""" % table._sequence_name) self.execute("select lastval()")
return int(self.cursor.fetchone()[0]) return int(self.cursor.fetchone()[0])
def try_json(self): def try_json(self):
@@ -2942,6 +2993,11 @@ class PostgreSQLAdapter(BaseAdapter):
return value return value
return BaseAdapter.represent(self, obj, fieldtype) return BaseAdapter.represent(self, obj, fieldtype)
def _drop(self, table, mode='restrict'):
if mode not in ['restrict', 'cascade', '']:
raise ValueError('Invalid mode: %s' % mode)
return ['DROP TABLE ' + table.sqlsafe + ' ' + str(mode) + ';']
class NewPostgreSQLAdapter(PostgreSQLAdapter): class NewPostgreSQLAdapter(PostgreSQLAdapter):
drivers = ('psycopg2','pg8000') drivers = ('psycopg2','pg8000')
@@ -3074,8 +3130,6 @@ class OracleAdapter(BaseAdapter):
'reference TFK': ' CONSTRAINT FK_%(foreign_table)s_PK FOREIGN KEY (%(field_name)s) REFERENCES %(foreign_table)s (%(foreign_key)s) ON DELETE %(on_delete_action)s', 'reference TFK': ' CONSTRAINT FK_%(foreign_table)s_PK FOREIGN KEY (%(field_name)s) REFERENCES %(foreign_table)s (%(foreign_key)s) ON DELETE %(on_delete_action)s',
} }
def sequence_name(self,tablename):
return '%s_sequence' % tablename
def trigger_name(self,tablename): def trigger_name(self,tablename):
return '%s_trigger' % tablename return '%s_trigger' % tablename
@@ -3091,7 +3145,7 @@ class OracleAdapter(BaseAdapter):
def _drop(self,table,mode): def _drop(self,table,mode):
sequence_name = table._sequence_name sequence_name = table._sequence_name
return ['DROP TABLE %s %s;' % (table, mode), 'DROP SEQUENCE %s;' % sequence_name] return ['DROP TABLE %s %s;' % (table.sqlsafe, mode), 'DROP SEQUENCE %s;' % sequence_name]
def select_limitby(self, sql_s, sql_f, sql_t, sql_w, sql_o, limitby): def select_limitby(self, sql_s, sql_f, sql_t, sql_w, sql_o, limitby):
if limitby: if limitby:
@@ -3218,6 +3272,13 @@ class OracleAdapter(BaseAdapter):
else: else:
return self.cursor.fetchall() return self.cursor.fetchall()
def sqlsafe_table(self, tablename, ot=None):
if ot is not None:
return (self.QUOTE_TEMPLATE + ' ' \
+ self.QUOTE_TEMPLATE) % (ot, tablename)
return self.QUOTE_TEMPLATE % tablename
class MSSQLAdapter(BaseAdapter): class MSSQLAdapter(BaseAdapter):
drivers = ('pyodbc',) drivers = ('pyodbc',)
T_SEP = 'T' T_SEP = 'T'
@@ -3685,7 +3746,7 @@ class FireBirdAdapter(BaseAdapter):
} }
def sequence_name(self,tablename): def sequence_name(self,tablename):
return 'genid_%s' % tablename return ('genid_' + self.QUOTE_TEMPLATE) % tablename
def trigger_name(self,tablename): def trigger_name(self,tablename):
return 'trg_id_%s' % tablename return 'trg_id_%s' % tablename
@@ -3714,7 +3775,7 @@ class FireBirdAdapter(BaseAdapter):
def _drop(self,table,mode): def _drop(self,table,mode):
sequence_name = table._sequence_name sequence_name = table._sequence_name
return ['DROP TABLE %s %s;' % (table, mode), 'DROP GENERATOR %s;' % sequence_name] return ['DROP TABLE %s %s;' % (table.sqlsafe, mode), 'DROP GENERATOR %s;' % sequence_name]
def select_limitby(self, sql_s, sql_f, sql_t, sql_w, sql_o, limitby): def select_limitby(self, sql_s, sql_f, sql_t, sql_w, sql_o, limitby):
if limitby: if limitby:
@@ -4283,7 +4344,7 @@ class SAPDBAdapter(BaseAdapter):
} }
def sequence_name(self,table): def sequence_name(self,table):
return '%s_id_Seq' % table return (self.QUOTE_TEMPLATE + '_id_Seq') % table
def select_limitby(self, sql_s, sql_f, sql_t, sql_w, sql_o, limitby): def select_limitby(self, sql_s, sql_f, sql_t, sql_w, sql_o, limitby):
if limitby: if limitby:
@@ -5446,8 +5507,8 @@ def cleanup(text):
""" """
validates that the given text is clean: only contains [0-9a-zA-Z_] validates that the given text is clean: only contains [0-9a-zA-Z_]
""" """
if not REGEX_ALPHANUMERIC.match(text): #if not REGEX_ALPHANUMERIC.match(text):
raise SyntaxError('invalid table or field name: %s' % text) # raise SyntaxError('invalid table or field name: %s' % text)
return text return text
class MongoDBAdapter(NoSQLAdapter): class MongoDBAdapter(NoSQLAdapter):
@@ -7249,12 +7310,32 @@ class Row(object):
__init__ = lambda self,*args,**kwargs: self.__dict__.update(*args,**kwargs) __init__ = lambda self,*args,**kwargs: self.__dict__.update(*args,**kwargs)
def __getitem__(self, k): def __getitem__(self, k):
if isinstance(k, Table):
try:
return ogetattr(self, k._tablename)
except (KeyError,AttributeError,TypeError):
pass
elif isinstance(k, Field):
try:
return ogetattr(self, k.name)
except (KeyError,AttributeError,TypeError):
pass
try:
return ogetattr(ogetattr(self, k.tablename), k.name)
except (KeyError,AttributeError,TypeError):
pass
key=str(k) key=str(k)
_extra = self.__dict__.get('_extra', None) _extra = ogetattr(self, '__dict__').get('_extra', None)
if _extra is not None: if _extra is not None:
v = _extra.get(key, DEFAULT) v = _extra.get(key, DEFAULT)
if v != DEFAULT: if v != DEFAULT:
return v return v
try:
return ogetattr(self, key)
except (KeyError,AttributeError,TypeError):
pass
m = REGEX_TABLE_DOT_FIELD.match(key) m = REGEX_TABLE_DOT_FIELD.match(key)
if m: if m:
try: try:
@@ -7656,7 +7737,7 @@ class DAL(object):
adapter_args=None, attempts=5, auto_import=False, adapter_args=None, attempts=5, auto_import=False,
bigint_id=False, debug=False, lazy_tables=False, bigint_id=False, debug=False, lazy_tables=False,
db_uid=None, do_connect=True, db_uid=None, do_connect=True,
after_connection=None, tables=None): after_connection=None, tables=None, ignore_field_case=True):
""" """
Creates a new Database Abstraction Layer instance. Creates a new Database Abstraction Layer instance.
@@ -7741,6 +7822,7 @@ class DAL(object):
self._decode_credentials = decode_credentials self._decode_credentials = decode_credentials
self._attempts = attempts self._attempts = attempts
self._do_connect = do_connect self._do_connect = do_connect
self._ignore_field_case = ignore_field_case
if not str(attempts).isdigit() or attempts < 0: if not str(attempts).isdigit() or attempts < 0:
attempts = 5 attempts = 5
@@ -7772,6 +7854,7 @@ class DAL(object):
# copy so multiple DAL() possible # copy so multiple DAL() possible
self._adapter.types = copy.copy(types) self._adapter.types = copy.copy(types)
self._adapter.build_parsemap() self._adapter.build_parsemap()
self._adapter.ignore_field_case = ignore_field_case
if bigint_id: if bigint_id:
if 'big-id' in types and 'reference' in types: if 'big-id' in types and 'reference' in types:
self._adapter.types['id'] = types['big-id'] self._adapter.types['id'] = types['big-id']
@@ -8569,13 +8652,8 @@ class Table(object):
self._tablename = tablename self._tablename = tablename
self._ot = None # args.get('rname') self._ot = None # args.get('rname')
self._rname = args.get('rname') self._rname = args.get('rname')
if not self._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(tablename)
else:
tb = self._rname[1:-1]
self._sequence_name = args.get('sequence_name') or \
db and db._adapter.sequence_name(tb)
self._trigger_name = args.get('trigger_name') or \ self._trigger_name = args.get('trigger_name') or \
db and db._adapter.trigger_name(tablename) db and db._adapter.trigger_name(tablename)
self._common_filter = args.get('common_filter') self._common_filter = args.get('common_filter')
@@ -8656,7 +8734,7 @@ class Table(object):
fields.append(Field(fn,'blob',default='', fields.append(Field(fn,'blob',default='',
writable=False,readable=False)) writable=False,readable=False))
lower_fieldnames = set() fieldnames_set = set()
reserved = dir(Table) + ['fields'] reserved = dir(Table) + ['fields']
if (db and db.check_reserved): if (db and db.check_reserved):
check_reserved = db.check_reserved_keyword check_reserved = db.check_reserved_keyword
@@ -8667,12 +8745,15 @@ class Table(object):
for field in fields: for field in fields:
field_name = field.name field_name = field.name
check_reserved(field_name) check_reserved(field_name)
fn_lower = field_name.lower() if db and db._ignore_field_case:
if fn_lower in lower_fieldnames: fname_item = field_name.lower()
else:
fname_item = field_name
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:
lower_fieldnames.add(fn_lower) fieldnames_set.add(fname_item)
self.fields.append(field_name) self.fields.append(field_name)
self[field_name] = field self[field_name] = field
@@ -8914,8 +8995,23 @@ class Table(object):
if 'Oracle' in str(type(self._db._adapter)): if 'Oracle' in str(type(self._db._adapter)):
return '%s %s' % (ot, self._tablename) return '%s %s' % (ot, self._tablename)
return '%s AS %s' % (ot, self._tablename) return '%s AS %s' % (ot, self._tablename)
return self._tablename return self._tablename
@property
def sqlsafe(self):
rname = self._rname
if rname: return rname
return self._db._adapter.sqlsafe_table(self._tablename)
@property
def sqlsafe_alias(self):
rname = self._rname
ot = self._ot
if rname and not ot: return rname
return self._db._adapter.sqlsafe_table(self._tablename, self._ot)
def _drop(self, mode = ''): def _drop(self, mode = ''):
return self._db._adapter._drop(self, mode) return self._db._adapter._drop(self, mode)
@@ -9742,7 +9838,12 @@ class Field(Expression):
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 REGEX_PYTHON_KEYWORDS.match(fieldname):
raise SyntaxError('Field: invalid field name: %s' % fieldname) raise SyntaxError('Field: invalid field name: %s' % fieldname)
self.type = type if not isinstance(type, (Table,Field)) else 'reference %s' % type
if not isinstance(type, (Table,Field)):
self.type = type
else:
self.type = 'reference %s' % type
self.length = length if not length is None else DEFAULTLENGTH.get(self.type,512) self.length = length if not length is None else DEFAULTLENGTH.get(self.type,512)
self.default = default if default!=DEFAULT else (update or None) self.default = default if default!=DEFAULT else (update or None)
self.required = required # is this field required self.required = required # is this field required
@@ -10008,6 +10109,16 @@ class Field(Expression):
except: except:
return '<no table>.%s' % self.name return '<no table>.%s' % self.name
@property
def sqlsafe(self):
if self._table:
return self._table.sqlsafe + '.' + (self._rname or self._db._adapter.sqlsafe_field(self.name))
return '<no table>.%s' % self.name
@property
def sqlsafe_name(self):
return self._rname or self._db._adapter.sqlsafe_field(self.name)
class Query(object): class Query(object):
@@ -10360,7 +10471,7 @@ class Set(object):
fields = table._listify(update_fields,update=True) fields = table._listify(update_fields,update=True)
if not fields: if not fields:
raise SyntaxError("No fields to update") raise SyntaxError("No fields to update")
ret = db._adapter.update("%s" % table,self.query,fields) ret = db._adapter.update("%s" % table._tablename,self.query,fields)
ret and [f(self,update_fields) for f in table._after_update] ret and [f(self,update_fields) for f in table._after_update]
return ret return ret
@@ -10873,11 +10984,21 @@ 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):
unq_colnames = []
for col in colnames:
m = self.db._adapter.REGEX_TABLE_DOT_FIELD.match(col)
if not m:
unq_colnames.append(col)
else:
unq_colnames.append('.'.join(m.groups()))
return unq_colnames
colnames = kwargs.get('colnames', self.colnames) colnames = kwargs.get('colnames', self.colnames)
write_colnames = kwargs.get('write_colnames',True) write_colnames = kwargs.get('write_colnames',True)
# a proper csv starting with the column names # a proper csv starting with the column names
if write_colnames: if write_colnames:
writer.writerow(colnames) writer.writerow(unquote_colnames(colnames))
def none_exception(value): def none_exception(value):
""" """
@@ -10900,10 +11021,11 @@ class Rows(object):
for record in self: for record in self:
row = [] row = []
for col in colnames: for col in colnames:
if not REGEX_TABLE_DOT_FIELD.match(col): m = self.db._adapter.REGEX_TABLE_DOT_FIELD.match(col)
if not m:
row.append(record._extra[col]) row.append(record._extra[col])
else: else:
(t, f) = col.split('.') (t, f) = m.groups()
field = self.db[t][f] field = self.db[t][f]
if isinstance(record.get(t, None), (Row,dict)): if isinstance(record.get(t, None), (Row,dict)):
value = record[t][f] value = record[t][f]
+3 -3
View File
@@ -21,7 +21,7 @@ from gluon.html import FORM, INPUT, LABEL, OPTION, SELECT, COL, COLGROUP
from gluon.html import TABLE, THEAD, TBODY, TR, TD, TH, STYLE from gluon.html import TABLE, THEAD, TBODY, TR, TD, TH, STYLE
from gluon.html import URL, truncate_string, FIELDSET from gluon.html import URL, truncate_string, FIELDSET
from gluon.dal import DAL, Field, Table, Row, CALLABLETYPES, smart_query, \ from gluon.dal import DAL, Field, Table, Row, CALLABLETYPES, smart_query, \
bar_encode, Reference, REGEX_TABLE_DOT_FIELD, Expression, SQLCustomType bar_encode, Reference, Expression, SQLCustomType
from gluon.storage import Storage from gluon.storage import Storage
from gluon.utils import md5_hash from gluon.utils import md5_hash
from gluon.validators import IS_EMPTY_OR, IS_NOT_EMPTY, IS_LIST_OF, IS_DATE, \ from gluon.validators import IS_EMPTY_OR, IS_NOT_EMPTY, IS_LIST_OF, IS_DATE, \
@@ -2893,7 +2893,7 @@ class SQLTABLE(TABLE):
if not sqlrows: if not sqlrows:
return return
if not columns: if not columns:
columns = sqlrows.colnames columns = ['.'.join(sqlrows.db._adapter.REGEX_TABLE_DOT_FIELD.match(c).groups()) for c in sqlrows.colnames]
if headers == 'fieldname:capitalize': if headers == 'fieldname:capitalize':
headers = {} headers = {}
for c in columns: for c in columns:
@@ -3116,7 +3116,7 @@ class ExportClass(object):
for record in self.rows: for record in self.rows:
row = [] row = []
for col in self.rows.colnames: for col in self.rows.colnames:
if not REGEX_TABLE_DOT_FIELD.match(col): if not self.rows.db._adapter.REGEX_TABLE_DOT_FIELD.match(col):
row.append(record._extra[col]) row.append(record._extra[col])
else: else:
(t, f) = col.split('.') (t, f) = col.split('.')
+112 -1
View File
@@ -79,6 +79,9 @@ 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')
@@ -1393,7 +1396,115 @@ 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):
# tests for complex table names
def testRun(self):
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)
t0.drop('cascade')
t1.drop()
db.t1.drop()
db.t0.drop()
# tests for case sensitivity
def testCase(self):
db = DAL(DEFAULT_URI, check_reserved=['all'], ignore_field_case=False)
# test table case
t0 = db.define_table('B',
Field('f', 'string'))
try:
t1 = db.define_table('b',
Field('B', t0),
Field('words', 'text'))
except Exception, e:
# An error is expected when database does not support case
# sensitive entity names.
if DEFAULT_URI.startswith('sqlite:'):
self.assertTrue(isinstance(e, db._adapter.driver.OperationalError))
return
raise e
blather = 'blah blah and so'
t0[0] = {'f': 'content'}
t1[0] = {'B': int(t0[1]['id']),
'words': blather}
r = db(db.B.id==db.b.B).select()
self.assertEqual(r[0].b.words, blather)
t1.drop()
t0.drop()
# test field case
try:
t0 = db.define_table('table is a test',
Field('a_a'),
Field('a_A'))
except Exception, e:
# some db does not support case sensitive field names mysql is one of them.
if DEFAULT_URI.startswith('mysql:'):
db.rollback()
return
raise e
t0[0] = dict(a_a = 'a_a', a_A='a_A')
self.assertEqual(t0[1].a_a, 'a_a')
self.assertEqual(t0[1].a_A, 'a_A')
t0.drop()
def testPKFK(self):
# test primary keys
db = DAL(DEFAULT_URI, check_reserved=['all'], ignore_field_case=False)
# test table without surrogate key. Length must is limited to
# 100 because of MySQL limitations: it cannot handle more than
# 767 bytes in unique keys.
t0 = db.define_table('t0', Field('Code', length=100), primarykey=['Code'])
t22 = db.define_table('t22', Field('f'), Field('t0_Code', 'reference t0'))
t3 = db.define_table('t3', Field('f', length=100), Field('t0_Code', t0.Code), primarykey=['f'])
t4 = db.define_table('t4', Field('f', length=100), Field('t0', t0), primarykey=['f'])
try:
t5 = db.define_table('t5', Field('f', length=100), Field('t0', 'reference no_table_wrong_reference'), primarykey=['f'])
except Exception, e:
self.assertTrue(isinstance(e, KeyError))
t0.drop('cascade')
t22.drop()
t3.drop()
t4.drop()
if __name__ == '__main__': if __name__ == '__main__':
unittest.main() unittest.main()
tearDownModule() tearDownModule()