Files
YMS/yms/model/dbrecord.py
T

150 lines
4.4 KiB
Python

# SPDX-License-Identifier: WTFPL
"""
No commits here: Has to be done elsewhere to support transactions!
"""
from datetime import datetime
class DbRecord:
def __init__(self, db, record_id=None):
if not hasattr(self, "table"):
raise TypeError("DbRecord subclass must define 'table'")
if not hasattr(self, "pk_field"):
raise TypeError("DbRecord subclass must define 'pk_field'")
if not hasattr(self, "fields"):
raise TypeError("DbRecord subclass must define 'fields'")
self.db = db
self.cur = db.cur
self.dcur = db.dcur
self.conn = db.conn
self.last_update = None
self.title = _("DbRecord({})") # for pretty print
if record_id is not None:
self.pk_value = int(record_id)
self.load()
else:
self.clear()
def clear(self):
self.pk_value = None
self.data = {} # for old/new checks
for field in self.fields:
setattr(self, field, None)
self.data[field] = None
def load(self, record_id=None):
if record_id is not None:
self.pk_value = int(record_id)
if self.pk_value is None:
self.clear()
return
sql = ("SELECT {} FROM {} WHERE {}=%s").format(', '.join(self.fields), self.table, self.pk_field)
try:
self.dcur.execute(sql, (self.pk_value,))
self.data = self.dcur.fetchone()
except Exception as e:
print(e)
print(sql)
self.db.error = e
self.clear()
return False
if self.data:
for field, value in self.data.items():
setattr(self, field, value)
return True
else:
self.clear()
return False
def save(self):
self.last_update = datetime.now()
if self.pk_value is not None:
return self._update()
else:
return self._insert()
def _update(self):
sqlfield = []
values = []
for field in self.fields:
newdata = getattr(self, field)
if self.data[field] != newdata:
sqlfield.append("{}=%s".format(field))
values.append(newdata)
if values:
# only save if there are changes
sql = (
f"UPDATE {self.table} SET " + ",".join(sqlfield) +
f" WHERE {self.pk_field}=%s"
)
values.append(self.pk_value)
try:
self.dcur.execute(sql, values)
except Exception as e:
print(e)
print(sql)
print(values)
self.db.error = e
return False
for field in self.fields:
self.data[field] = getattr(self, field)
return True
def _insert(self):
sql = (
"INSERT INTO {} ({}) "
"VALUES ({})"
).format(
self.table,
','.join(self.fields),
','.join(["%s"] * len(self.fields))
)
values = []
for field in self.fields:
newdata = getattr(self, field)
values.append(newdata)
self.data[field] = newdata
try:
self.dcur.execute(sql, values)
except Exception as e:
print(e)
print(sql)
print(values)
self.db.error = e
return False
self.pk_value = self.dcur.lastrowid
return True
def delete(self, record_id=None):
if record_id is None:
if self.pk_value is None:
return False
record_id = self.pk_value
sql = "DELETE FROM {} WHERE {}=%s".format(self.table, self.pk_field)
try:
self.dcur.execute(sql, (record_id,))
except Exception as e:
print(e)
print(sql)
self.db.error = e
return False
if self.dcur.rowcount != 1:
return False
if self.pk_value == record_id:
self.pk_value = None
self.clear()
return True
def __str__(self):
out = [self.title.format(self.pk_value)]
for field in self.fields:
value = getattr(self, field, None)
out.append(" {:15}: {}".format(field, value))
return "\n".join(out)