# 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.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.pk_value = None self.clear() def clear(self): for field in self.fields: setattr(self, field, None) self.data = {} # for old/new checks 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: 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: 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: 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: 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)