quoting of tablenames and field names

fix for wrong sequence name quoting

fix for wrong sequence name quoting 2

reverted test to quoted rname

rname passed asis. freedom in table names and field names. Case sensitivity in field and table names

removed test of allowed char sequence in table and field names, there is no limit anymore

fixed some troubles with postgresql drop(). Better Row.__getitem__(): should not be much slower.  Added a ignore_field_case parameter to allow Field('F')!=Field('f') on databases supporting case sesitivity.

different mysql drivers handled in quoting test.

improved handling of foreing keys

fix with quoted names in SQLTABLE.  fix for ORDER BY

handle references on non "id" fields

check reference expression against table list first
This commit is contained in:
Michele Comitini
2013-11-08 17:11:36 +01:00
parent 645c3b10dd
commit 14bcad629e
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()