169 lines
5.4 KiB
Python
169 lines
5.4 KiB
Python
import MySQLdb
|
|
import secrets
|
|
|
|
class Database:
|
|
|
|
def __init__(self, hostname, dbname):
|
|
self.hostname = hostname
|
|
self.schema = dbname
|
|
self.conn = None
|
|
self.cur = None
|
|
self.dcur = None
|
|
self.error = None
|
|
|
|
def connect(self, dbuser, dbpass):
|
|
if self.conn:
|
|
return True
|
|
MySQLdb.paramstyle = 'pyformat'
|
|
try:
|
|
self.conn = MySQLdb.connect(host=self.hostname, db=self.schema,
|
|
user=dbuser, passwd=dbpass, charset='utf8')
|
|
except MySQLdb.Error as e:
|
|
self.error = e
|
|
return False
|
|
self.conn.autocommit(False)
|
|
self.cur = self.conn.cursor()
|
|
self.dcur = self.conn.cursor(MySQLdb.cursors.DictCursor)
|
|
return True
|
|
|
|
def connect_local(self):
|
|
# connection via local unix socket
|
|
if self.conn:
|
|
return True
|
|
MySQLdb.paramstyle = 'pyformat'
|
|
try:
|
|
self.conn = MySQLdb.connect(host='localhost', db=self.schema,
|
|
unix_socket='/run/mysqld/mysqld.sock', charset='utf8')
|
|
except MySQLdb.Error as e:
|
|
self.error = e
|
|
return False
|
|
self.cur = self.conn.cursor()
|
|
self.dcur = self.conn.cursor(MySQLdb.cursors.DictCursor)
|
|
return True
|
|
|
|
def disconnect(self):
|
|
self.cur.close()
|
|
self.dcur.close()
|
|
self.conn.close()
|
|
|
|
def is_connected(self):
|
|
return self.conn.open
|
|
|
|
def get_timestamp(self):
|
|
self.cur.execute("SELECT CURRENT_TIMESTAMP")
|
|
return self.cur.fetchone()[0]
|
|
|
|
def last_error(self):
|
|
if not self.error:
|
|
return(0, '')
|
|
code = self.error.args[0]
|
|
text = self.error.args[1]
|
|
self.error = None
|
|
return(code, text)
|
|
|
|
def load_enum(self, table, column):
|
|
"""
|
|
TODO this is only an example! perhaps other keys needed
|
|
Load an enum directly from a database table and provide it as a
|
|
dictionary for use in, for example, selection lists.
|
|
"""
|
|
sql = (
|
|
"SELECT TRIM(TRAILING ')' FROM SUBSTRING(column_type,6)) "
|
|
"FROM information_schema.columns "
|
|
"WHERE table_schema=DATABASE() AND table_name=%s AND column_name=%s AND data_type='enum'"
|
|
)
|
|
self.cur.execute(sql, (table, column))
|
|
row = self.cur.fetchone()
|
|
if not row:
|
|
return {}
|
|
return {
|
|
i: _(value.strip("'"))
|
|
for i, value in enumerate(row[0].split(','), 1)
|
|
}
|
|
|
|
def lock_acquire(self, table, key, userid):
|
|
token = secrets.token_hex(8)
|
|
# first try new lock
|
|
sql = (
|
|
"INSERT INTO recordlock (lock_table, lock_key, token, userid)"
|
|
"VALUES (%s, %s, %s, %s)"
|
|
)
|
|
try:
|
|
self.cur.execute(sql, (table, key, token, userid))
|
|
self.conn.commit()
|
|
return token
|
|
except self.conn.IntegrityError as e:
|
|
self.conn.rollback()
|
|
if e.errno != 1062:
|
|
raise
|
|
# lazy locking: lock exists but is old, so we use it
|
|
sql = (
|
|
"UPDATE recordlock "
|
|
"SET token=%s, userid=%s,"
|
|
" created_at=CURRENT_TIMESTAMP, last_seen=CURRENT_TIMESTAMP "
|
|
"WHERE lock_table=%s AND lock_key=%s AND last_seen<CURRENT_TIMESTAMP-INTERVAL 5 MINUTE"
|
|
)
|
|
self.cur.execute(sql, (token, userid, table, key))
|
|
self.conn.commit()
|
|
if self.cur.rowcount == 1:
|
|
return token
|
|
return None
|
|
|
|
def lock_refresh(self, token):
|
|
sql = (
|
|
"UPDATE recordlock "
|
|
"SET last_seen=current_timestamp() "
|
|
"WHERE token=%s"
|
|
)
|
|
self.cur.execute(sql, (token,))
|
|
self.conn.commit()
|
|
return self.cur.rowcount == 1
|
|
|
|
def lock_release(self, token):
|
|
sql = "DELETE FROM recordlock WHERE token=%s"
|
|
self.cur.execute(sql, (token,))
|
|
self.conn.commit()
|
|
return self.cur.rowcount == 1
|
|
|
|
def load_setting(self, sno):
|
|
sql = "SELECT valstr, valint FROM settings WHERE userid=0 AND sno=%s"
|
|
self.cur.execute(sql, (sno,))
|
|
row = self.cur.fetchone()
|
|
if not row:
|
|
sql = "INSERT INTO settings (userid, sno) VALUES (0, %s)"
|
|
self.cur.execute(sql, (sno,))
|
|
self.conn.commit()
|
|
return (None, None)
|
|
return row
|
|
|
|
def save_setting(self, sno, valstr, valint):
|
|
sql = "UPDATE settings SET valstr=%s, valint=%s WHERE userid=0 AND sno=%s"
|
|
values = (valstr, valint, sno)
|
|
return self.execute_write(sql, values, 1)
|
|
|
|
def execute_write(self, sql, values=(), expected_rows=None):
|
|
"""
|
|
Execute an INSERT, UPDATE, or DELETE statement.
|
|
expected_rows:
|
|
None = do not check the number of affected rows
|
|
1 = expect exactly one affected row
|
|
0 = expect no affected rows
|
|
"""
|
|
try:
|
|
self.cur.execute(sql, values)
|
|
except MySQLdb.Error as e:
|
|
self.rollback()
|
|
self.error = e
|
|
return False
|
|
if (expected_rows is not None and self.cur.rowcount != expected_rows):
|
|
self.rollback()
|
|
self.error = RuntimeError(
|
|
"Expected {} affected rows, got {}".format(
|
|
expected_rows,
|
|
self.cur.rowcount
|
|
)
|
|
)
|
|
return False
|
|
self.error = None
|
|
return True
|