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_DBNAME = re.compile('^(\w+)(\:\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_CLEANUP_FN = re.compile('[\'"\s;]+')
REGEX_UNPACK = re.compile('(?<!\|)\|(?!\|)')
@@ -670,12 +671,27 @@ class ConnectionPool(object):
break
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
###################################################################################
class BaseAdapter(ConnectionPool):
__metaclass__ = AdapterMeta
native_json = False
driver = None
driver_name = None
@@ -693,6 +709,7 @@ class BaseAdapter(ConnectionPool):
T_SEP = ' '
QUOTE_TEMPLATE = '"%s"'
types = {
'boolean': 'CHAR(1)',
'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'
'big-id': 'BIGINT PRIMARY KEY AUTOINCREMENT',
'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):
@@ -834,8 +852,9 @@ class BaseAdapter(ConnectionPool):
self.connection = Dummy()
self.cursor = Dummy()
def sequence_name(self,tablename):
return '%s_sequence' % tablename
return self.QUOTE_TEMPLATE % ('%s_sequence' % tablename)
def trigger_name(self,tablename):
return '%s_sequence' % tablename
@@ -868,61 +887,72 @@ class BaseAdapter(ConnectionPool):
if referenced == '.':
referenced = tablename
constraint_name = self.constraint_name(tablename, field_name)
if not '.' in referenced \
and referenced != tablename \
and hasattr(table,'_primarykey'):
ftype = types['integer']
else:
if hasattr(table,'_primarykey'):
# if not '.' in referenced \
# and referenced != tablename \
# and hasattr(table,'_primarykey'):
# ftype = types['integer']
#else:
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('.')
rtable = db[rtablename]
rfield = rtable[rfieldname]
# must be PK reference or unique
if rfieldname in rtable._primarykey or \
rfield.unique:
ftype = types[rfield.type[:9]] % \
dict(length=rfield.length)
# multicolumn primary key reference?
if not rfield.unique and len(rtable._primarykey)>1:
# then it has to be a table level FK
if rtablename not in TFK:
TFK[rtablename] = {}
TFK[rtablename][rfieldname] = field_name
else:
ftype = ftype + \
types['reference FK'] % dict(
constraint_name = constraint_name, # should be quoted
foreign_key = '%s (%s)' % (rtablename,
rfieldname),
table_name = tablename,
field_name = field._rname or field.name,
on_delete_action=field.ondelete)
except Exception, e:
LOGGER.debug('Error: %s' %e)
raise KeyError('Cannot resolve reference %s in %s definition' % (referenced, table._tablename))
# must be PK reference or unique
if getattr(rtable, '_primarykey', None) and rfieldname in rtable._primarykey or \
rfield.unique:
ftype = types[rfield.type[:9]] % \
dict(length=rfield.length)
# multicolumn primary key reference?
if not rfield.unique and len(rtable._primarykey)>1:
# then it has to be a table level FK
if rtablename not in TFK:
TFK[rtablename] = {}
TFK[rtablename][rfieldname] = field_name
else:
# make a guess here for circular references
if referenced in db:
id_fieldname = db[referenced]._id.name
elif referenced == tablename:
id_fieldname = table._id.name
else: #make a guess
id_fieldname = 'id'
#gotcha: the referenced table must be defined before
#the referencing one to be able to create the table
#Also if it's not recommended, we can still support
#references to tablenames without rname to make
#migrations and model relationship work also if tables
#are not defined in order
real_referenced = (
(db[referenced]._rname or db[referenced])
if referenced == tablename or referenced in db
else referenced)
ftype = types[field_type[:9]] % dict(
index_name = field_name+'__idx',
field_name = field._rname or field.name,
constraint_name = constraint_name,
foreign_key = '%s (%s)' % (real_referenced,
id_fieldname),
on_delete_action=field.ondelete)
ftype = ftype + \
types['reference FK'] % dict(
constraint_name = constraint_name, # should be quoted
foreign_key = rtable.sqlsafe + ' (' + rfield.sqlsafe_name + ')',
table_name = table.sqlsafe,
field_name = field.sqlsafe_name,
on_delete_action=field.ondelete)
else:
# make a guess here for circular references
if referenced in db:
id_fieldname = db[referenced]._id.sqlsafe_name
elif referenced == tablename:
id_fieldname = table._id.sqlsafe_name
else: #make a guess
id_fieldname = self.QUOTE_TEMPLATE % 'id'
#gotcha: the referenced table must be defined before
#the referencing one to be able to create the table
#Also if it's not recommended, we can still support
#references to tablenames without rname to make
#migrations and model relationship work also if tables
#are not defined in order
if referenced == tablename:
real_referenced = db[referenced].sqlsafe
else:
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'):
ftype = types[field_type[:14]]
elif field_type.startswith('decimal'):
@@ -995,39 +1025,37 @@ class BaseAdapter(ConnectionPool):
# geometry fields are added after the table has been created, not now
if not (self.dbengine == 'postgres' and \
field_type.startswith('geom')):
#fetch the rname if it's there
field_rname = "%s" % (field._rname or field_name)
fields.append('%s %s' % (field_rname, ftype))
fields.append('%s %s' % (field.sqlsafe_name, ftype))
other = ';'
# backend-specific extensions to fields
if self.dbengine == 'mysql':
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;'
fields = ',\n '.join(fields)
for rtablename in TFK:
rfields = TFK[rtablename]
pkeys = db[rtablename]._primarykey
fkeys = [ rfields[k] for k in pkeys ]
pkeys = [self.QUOTE_TEMPLATE % pk for pk in db[rtablename]._primarykey]
fkeys = [self.QUOTE_TEMPLATE % rfields[k].name for k in pkeys ]
fields = fields + ',\n ' + \
types['reference TFK'] % dict(
table_name = tablename,
table_name = table.sqlsafe,
field_name=', '.join(fkeys),
foreign_table = rtablename,
foreign_table = table.sqlsafe,
foreign_key = ', '.join(pkeys),
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):
query = "CREATE TABLE %s(\n %s,\n %s) %s" % \
(table_rname, fields,
self.PRIMARY_KEY(', '.join(table._primarykey)),other)
(table.sqlsafe, fields,
self.PRIMARY_KEY(', '.join([self.QUOTE_TEMPLATE % pk for pk in table._primarykey])),other)
else:
query = "CREATE TABLE %s(\n %s\n)%s" % \
(table_rname, fields, other)
(table.sqlsafe, fields, other)
if self.uri.startswith('sqlite:///') \
or self.uri.startswith('spatialite:///'):
@@ -1105,6 +1133,7 @@ class BaseAdapter(ConnectionPool):
k,v=item
if not isinstance(v,dict):
v=dict(type='unknown',sql=v)
if self.ignore_field_case is not True: return k, v
return k.lower(),v
# make sure all field names are lower case to avoid
# migrations because of case cahnge
@@ -1132,7 +1161,7 @@ class BaseAdapter(ConnectionPool):
query = [ sql_fields[key]['sql'] ]
else:
query = ['ALTER TABLE %s ADD %s %s;' % \
(tablename, key,
(table.sqlsafe, key,
sql_fields_aux[key]['sql'].replace(', ', new_add))]
metadata_change = True
elif self.dbengine in ('sqlite', 'spatialite'):
@@ -1150,10 +1179,11 @@ class BaseAdapter(ConnectionPool):
"'%(table)s', '%(field)s');" %
dict(schema=schema, table=tablename, field=key,) ]
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:
query = ['ALTER TABLE %s DROP COLUMN %s;' %
(tablename, key)]
(self.QUOTE_TEMPLATE % tablename, self.QUOTE_TEMPLATE % key)]
metadata_change = True
elif sql_fields[key]['sql'] != sql_fields_old[key]['sql'] \
and not (key in table.fields and
@@ -1169,12 +1199,15 @@ class BaseAdapter(ConnectionPool):
else:
drop_expr = 'ALTER TABLE %s DROP COLUMN %s;'
key_tmp = key + '__tmp'
query = ['ALTER TABLE %s ADD %s %s;' % (t, key_tmp, tt),
'UPDATE %s SET %s=%s;' % (t, key_tmp, key),
drop_expr % (t, key),
'ALTER TABLE %s ADD %s %s;' % (t, key, tt),
'UPDATE %s SET %s=%s;' % (t, key, key_tmp),
drop_expr % (t, key_tmp)]
query = ['ALTER TABLE %s ADD %s %s;' % (self.QUOTE_TEMPLATE % t, self.QUOTE_TEMPLATE % key_tmp, tt),
'UPDATE %s SET %s=%s;' %
(self.QUOTE_TEMPLATE % t, self.QUOTE_TEMPLATE % key_tmp, self.QUOTE_TEMPLATE % key),
drop_expr % (self.QUOTE_TEMPLATE % t, self.QUOTE_TEMPLATE % key),
'ALTER TABLE %s ADD %s %s;' %
(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
elif sql_fields[key]['type'] != sql_fields_old[key]['type']:
sql_fields_current[key] = sql_fields[key]
@@ -1269,8 +1302,7 @@ class BaseAdapter(ConnectionPool):
return 'PRIMARY KEY(%s)' % key
def _drop(self, table, mode):
table_rname = table._rname or table
return ['DROP TABLE %s;' % table_rname]
return ['DROP TABLE %s;' % table.sqlsafe]
def drop(self, table, mode=''):
db = table._db
@@ -1288,17 +1320,16 @@ class BaseAdapter(ConnectionPool):
self.log('success!\n', table)
def _insert(self, table, fields):
table_rname = table._rname or table
table_rname = table.sqlsafe
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)
return 'INSERT INTO %s(%s) VALUES (%s);' % (table_rname, keys, values)
else:
return self._insert_empty(table)
def _insert_empty(self, table):
table_rname = table._rname or table
return 'INSERT INTO %s DEFAULT VALUES;' % table_rname
return 'INSERT INTO %s DEFAULT VALUES;' % (table.sqlsafe)
def insert(self, table, fields):
query = self._insert(table,fields)
@@ -1451,13 +1482,13 @@ class BaseAdapter(ConnectionPool):
self.expand(second, first.type))
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):
table_rname = first._ot and first or first._rname or first._tablename
table_rname = self.table_alias(first)
if use_common_filters(second):
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):
return '%s DESC' % self.expand(first)
@@ -1472,13 +1503,10 @@ class BaseAdapter(ConnectionPool):
if isinstance(expression, Field):
et = expression.table
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:
table_rname = et._tablename
if not colnames:
out = '%s.%s' % (table_rname, expression._rname or expression.name)
else:
out = '%s.%s' % (table_rname, expression.name)
out = '%s.%s' % (self.QUOTE_TEMPLATE % et._tablename, self.QUOTE_TEMPLATE % expression.name)
if field_type == 'string' and not expression.type in (
'string','text','json','password'):
out = self.CAST(out, self.types['text'])
@@ -1509,10 +1537,11 @@ class BaseAdapter(ConnectionPool):
else:
return str(expression)
def table_alias(self,name):
if not isinstance(name, Table):
name = self.db[name]._rname or self.db[name]
return str(name)
def table_alias(self, tbl):
if not isinstance(tbl, Table):
tbl = self.db[tbl]
return tbl.sqlsafe_alias
def alias(self, table, alias):
"""
@@ -1520,7 +1549,7 @@ class BaseAdapter(ConnectionPool):
with alias name.
"""
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['_tablename'] = alias
for fieldname in other.fields:
@@ -1532,8 +1561,7 @@ class BaseAdapter(ConnectionPool):
return other
def _truncate(self, table, mode=''):
tablename = table._rname or table._tablename
return ['TRUNCATE TABLE %s %s;' % (tablename, mode or '')]
return ['TRUNCATE TABLE %s %s;' % (table.sqlsafe, mode or '')]
def truncate(self, table, mode= ' '):
# Prepare functions "write_to_logfile" and "close_logfile"
@@ -1553,10 +1581,10 @@ class BaseAdapter(ConnectionPool):
sql_w = ' WHERE ' + self.expand(query)
else:
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)) \
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)
def update(self, tablename, query, fields):
@@ -1581,7 +1609,7 @@ class BaseAdapter(ConnectionPool):
sql_w = ' WHERE ' + self.expand(query)
else:
sql_w = ''
tablename = '%s' % (self.db[tablename]._rname or tablename)
tablename = self.db[tablename].sqlsafe
return 'DELETE FROM %s%s;' % (tablename, sql_w)
def delete(self, tablename, query):
@@ -1623,8 +1651,9 @@ class BaseAdapter(ConnectionPool):
if isinstance(item,SQLALL):
new_fields += item._table
elif isinstance(item,str):
if REGEX_TABLE_DOT_FIELD.match(item):
tablename,fieldname = item.split('.')
m = self.REGEX_TABLE_DOT_FIELD.match(item)
if m:
tablename,fieldname = m.groups()
append(db[tablename][fieldname])
else:
append(Expression(db,lambda item=item:item))
@@ -1645,10 +1674,11 @@ class BaseAdapter(ConnectionPool):
tablenames = tables(query)
tablenames_for_common_filters = tablenames
for field in fields:
if isinstance(field, basestring) \
and REGEX_TABLE_DOT_FIELD.match(field):
tn,fn = field.split('.')
field = self.db[tn][fn]
if isinstance(field, basestring):
m = self.REGEX_TABLE_DOT_FIELD.match(field)
if m:
tn,fn = m.groups()
field = self.db[tn][fn]
for tablename in tables(field):
if not tablename in tablenames:
tablenames.append(tablename)
@@ -1732,7 +1762,7 @@ class BaseAdapter(ConnectionPool):
tables_to_merge.keys()])
if joint:
sql_t += ' %s %s' % (command,
','.join([self.table_alias(t) for t in joint]))
','.join([t for t in joint]))
for t in joinon:
sql_t += ' %s %s' % (command, t)
elif inner_join and left:
@@ -1747,7 +1777,7 @@ class BaseAdapter(ConnectionPool):
sql_t += ' %s %s' % (icommand, t)
if joint:
sql_t += ' %s %s' % (command,
','.join([self.table_alias(t) for t in joint]))
','.join([t for t in joint]))
for t in joinon:
sql_t += ' %s %s' % (command, t)
else:
@@ -1767,9 +1797,9 @@ class BaseAdapter(ConnectionPool):
sql_o += ' ORDER BY %s' % self.expand(orderby)
if (limitby and not groupby and tablenames and orderby_on_limitby and not orderby):
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
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):
self.execute(query)
def log_execute(self, *a, **b):
if not self.connection: raise ValueError(a[0])
if not self.connection: return None
command = a[0]
if hasattr(self,'filter_sql_command'):
@@ -1955,8 +1987,17 @@ class BaseAdapter(ConnectionPool):
if field_is_type('decimal'):
return str(obj)
elif field_is_type('reference'): # reference
if fieldtype.find('.')>0:
return repr(obj)
# check for tablename first
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)):
return str(obj['id'])
return str(long(obj))
@@ -2162,14 +2203,15 @@ class BaseAdapter(ConnectionPool):
new_rows = []
tmps = []
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)
else:
(tablename, _the_sep_, fieldname) = colname.partition('.')
tablename, fieldname = col_m.groups()
table = db[tablename]
field = table[fieldname]
ft = field.type
tmps.append((tablename,fieldname,table,field,ft))
tmps.append((tablename, fieldname, table, field, ft))
for (i,row) in enumerate(rows):
new_row = Row()
for (j,colname) in enumerate(colnames):
@@ -2177,9 +2219,8 @@ class BaseAdapter(ConnectionPool):
tmp = tmps[j]
if tmp:
(tablename,fieldname,table,field,ft) = tmp
if tablename in new_row:
colset = new_row[tablename]
else:
colset = new_row.get(tablename, None)
if colset is None:
colset = new_row[tablename] = Row()
if tablename not in virtualtables:
virtualtables.append(tablename)
@@ -2287,6 +2328,14 @@ class BaseAdapter(ConnectionPool):
return Expression(self.db,'CASE WHEN %s THEN %s ELSE %s END' % \
(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.
###################################################################################
@@ -2564,6 +2613,7 @@ class MySQLAdapter(BaseAdapter):
'list:reference': 'LONGTEXT',
'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',
'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`"
@@ -2590,12 +2640,12 @@ class MySQLAdapter(BaseAdapter):
def _drop(self,table,mode):
# 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,
'SET FOREIGN_KEY_CHECKS=1;']
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):
self.execute('XA START;')
@@ -2668,6 +2718,8 @@ class MySQLAdapter(BaseAdapter):
class PostgreSQLAdapter(BaseAdapter):
drivers = ('psycopg2','pg8000')
QUOTE_TEMPLATE = '"%s"'
support_distributed_transaction = True
types = {
'boolean': 'CHAR(1)',
@@ -2694,12 +2746,11 @@ class PostgreSQLAdapter(BaseAdapter):
'geography': 'GEOGRAPHY',
'big-id': 'BIGSERIAL PRIMARY KEY',
'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 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 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',
}
QUOTE_TEMPLATE = '"%s"'
def varquote(self,name):
return varquote_aux(name,'"%s"')
@@ -2713,7 +2764,7 @@ class PostgreSQLAdapter(BaseAdapter):
return "'%s'" % str(obj).replace("'","''")
def sequence_name(self,table):
return '%s_id_seq' % table
return self.QUOTE_TEMPLATE % (table + '_id_seq')
def RANDOM(self):
return 'RANDOM()'
@@ -2802,8 +2853,8 @@ class PostgreSQLAdapter(BaseAdapter):
self.execute("SET standard_conforming_strings=on;")
self.try_json()
def lastrowid(self,table):
self.execute("""select currval('"%s"')""" % table._sequence_name)
def lastrowid(self,table = None):
self.execute("select lastval()")
return int(self.cursor.fetchone()[0])
def try_json(self):
@@ -2942,6 +2993,11 @@ class PostgreSQLAdapter(BaseAdapter):
return value
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):
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',
}
def sequence_name(self,tablename):
return '%s_sequence' % tablename
def trigger_name(self,tablename):
return '%s_trigger' % tablename
@@ -3091,7 +3145,7 @@ class OracleAdapter(BaseAdapter):
def _drop(self,table,mode):
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):
if limitby:
@@ -3218,6 +3272,13 @@ class OracleAdapter(BaseAdapter):
else:
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):
drivers = ('pyodbc',)
T_SEP = 'T'
@@ -3685,7 +3746,7 @@ class FireBirdAdapter(BaseAdapter):
}
def sequence_name(self,tablename):
return 'genid_%s' % tablename
return ('genid_' + self.QUOTE_TEMPLATE) % tablename
def trigger_name(self,tablename):
return 'trg_id_%s' % tablename
@@ -3714,7 +3775,7 @@ class FireBirdAdapter(BaseAdapter):
def _drop(self,table,mode):
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):
if limitby:
@@ -4283,7 +4344,7 @@ class SAPDBAdapter(BaseAdapter):
}
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):
if limitby:
@@ -5446,8 +5507,8 @@ def cleanup(text):
"""
validates that the given text is clean: only contains [0-9a-zA-Z_]
"""
if not REGEX_ALPHANUMERIC.match(text):
raise SyntaxError('invalid table or field name: %s' % text)
#if not REGEX_ALPHANUMERIC.match(text):
# raise SyntaxError('invalid table or field name: %s' % text)
return text
class MongoDBAdapter(NoSQLAdapter):
@@ -7249,12 +7310,32 @@ class Row(object):
__init__ = lambda self,*args,**kwargs: self.__dict__.update(*args,**kwargs)
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)
_extra = self.__dict__.get('_extra', None)
_extra = ogetattr(self, '__dict__').get('_extra', None)
if _extra is not None:
v = _extra.get(key, DEFAULT)
if v != DEFAULT:
return v
try:
return ogetattr(self, key)
except (KeyError,AttributeError,TypeError):
pass
m = REGEX_TABLE_DOT_FIELD.match(key)
if m:
try:
@@ -7656,7 +7737,7 @@ class DAL(object):
adapter_args=None, attempts=5, auto_import=False,
bigint_id=False, debug=False, lazy_tables=False,
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.
@@ -7741,6 +7822,7 @@ class DAL(object):
self._decode_credentials = decode_credentials
self._attempts = attempts
self._do_connect = do_connect
self._ignore_field_case = ignore_field_case
if not str(attempts).isdigit() or attempts < 0:
attempts = 5
@@ -7772,6 +7854,7 @@ class DAL(object):
# copy so multiple DAL() possible
self._adapter.types = copy.copy(types)
self._adapter.build_parsemap()
self._adapter.ignore_field_case = ignore_field_case
if bigint_id:
if 'big-id' in types and 'reference' in types:
self._adapter.types['id'] = types['big-id']
@@ -8569,13 +8652,8 @@ class Table(object):
self._tablename = tablename
self._ot = None # args.get('rname')
self._rname = args.get('rname')
if not self._rname:
self._sequence_name = args.get('sequence_name') or \
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._sequence_name = args.get('sequence_name') or \
db and db._adapter.sequence_name(self._rname or tablename)
self._trigger_name = args.get('trigger_name') or \
db and db._adapter.trigger_name(tablename)
self._common_filter = args.get('common_filter')
@@ -8656,7 +8734,7 @@ class Table(object):
fields.append(Field(fn,'blob',default='',
writable=False,readable=False))
lower_fieldnames = set()
fieldnames_set = set()
reserved = dir(Table) + ['fields']
if (db and db.check_reserved):
check_reserved = db.check_reserved_keyword
@@ -8667,12 +8745,15 @@ class Table(object):
for field in fields:
field_name = field.name
check_reserved(field_name)
fn_lower = field_name.lower()
if fn_lower in lower_fieldnames:
if db and db._ignore_field_case:
fname_item = field_name.lower()
else:
fname_item = field_name
if fname_item in fieldnames_set:
raise SyntaxError("duplicate field %s in table %s" \
% (field_name, tablename))
else:
lower_fieldnames.add(fn_lower)
fieldnames_set.add(fname_item)
self.fields.append(field_name)
self[field_name] = field
@@ -8914,8 +8995,23 @@ class Table(object):
if 'Oracle' in str(type(self._db._adapter)):
return '%s %s' % (ot, self._tablename)
return '%s AS %s' % (ot, 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 = ''):
return self._db._adapter._drop(self, mode)
@@ -9742,7 +9838,12 @@ class Field(Expression):
if not isinstance(fieldname, str) or hasattr(Table, fieldname) or \
fieldname[0] == '_' or REGEX_PYTHON_KEYWORDS.match(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.default = default if default!=DEFAULT else (update or None)
self.required = required # is this field required
@@ -10008,6 +10109,16 @@ class Field(Expression):
except:
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):
@@ -10360,7 +10471,7 @@ class Set(object):
fields = table._listify(update_fields,update=True)
if not fields:
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]
return ret
@@ -10873,11 +10984,21 @@ class Rows(object):
represent = kwargs.get('represent', False)
writer = csv.writer(ofile, delimiter=delimiter,
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)
write_colnames = kwargs.get('write_colnames',True)
# a proper csv starting with the column names
if write_colnames:
writer.writerow(colnames)
writer.writerow(unquote_colnames(colnames))
def none_exception(value):
"""
@@ -10900,10 +11021,11 @@ class Rows(object):
for record in self:
row = []
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])
else:
(t, f) = col.split('.')
(t, f) = m.groups()
field = self.db[t][f]
if isinstance(record.get(t, None), (Row,dict)):
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 URL, truncate_string, FIELDSET
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.utils import md5_hash
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:
return
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':
headers = {}
for c in columns:
@@ -3116,7 +3116,7 @@ class ExportClass(object):
for record in self.rows:
row = []
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])
else:
(t, f) = col.split('.')
+112 -1
View File
@@ -79,6 +79,9 @@ def tearDownModule():
class TestFields(unittest.TestCase):
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
self.assertRaises(SyntaxError, Field, '_abc', 'string')
@@ -1393,7 +1396,115 @@ class TestRNameFields(unittest.TestCase):
self.assertEqual(len(db.person._referenced_by),0)
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__':
unittest.main()
tearDownModule()
tearDownModule()