diff --git a/active_alchemy.py b/active_alchemy.py index f104c01..db75969 100644 --- a/active_alchemy.py +++ b/active_alchemy.py @@ -34,19 +34,22 @@ utcnow = arrow.utcnow + def _create_scoped_session(db, query_cls): - session = sessionmaker(autoflush=True, autocommit=False, - bind=db.engine, query_cls=query_cls) + session = sessionmaker( + autoflush=True, autocommit=False, bind=db.engine, query_cls=query_cls + ) return scoped_session(session) + def _tablemaker(db): def make_sa_table(*args, **kwargs): if len(args) > 1 and isinstance(args[1], db.Column): args = (args[0], db.metadata) + args[1:] - kwargs.setdefault('bind_key', None) - info = kwargs.pop('info', None) or {} - info.setdefault('bind_key', None) - kwargs['info'] = info + kwargs.setdefault("bind_key", None) + info = kwargs.pop("info", None) or {} + info.setdefault("bind_key", None) + kwargs["info"] = info return sqlalchemy.Table(*args, **kwargs) return make_sa_table @@ -103,11 +106,12 @@ class ModelTableNameDescriptor(object): """ Create the table name if it doesn't exist. """ + def __get__(self, obj, type): - tablename = type.__dict__.get('__tablename__') + tablename = type.__dict__.get("__tablename__") if not tablename: tablename = inflection.underscore(type.__name__) - setattr(type, '__tablename__', tablename) + setattr(type, "__tablename__", tablename) return tablename @@ -124,7 +128,7 @@ def get_engine(self): uri = self._sa_obj.uri info = self._sa_obj.info options = self._sa_obj.options - echo = options.get('echo') + echo = options.get("echo") if (uri, echo) == self._connected_for: return self._engine self._engine = engine = sqlalchemy.create_engine(info, **options) @@ -145,11 +149,11 @@ def __iter__(self): so we can do dict(sa_instance). """ for k in self.__dict__.keys(): - if not k.startswith('_'): + if not k.startswith("_"): yield (k, getattr(self, k)) def __repr__(self): - return '<%s>' % self.__class__.__name__ + return "<%s>" % self.__class__.__name__ def to_dict(self): """ @@ -187,6 +191,16 @@ def create(cls, **kwargs): record = cls(**kwargs).save() return record + @classmethod + def create_from_inst(cls, inst: object): + """ + Creates a new record from the params of a pre-existing object + :returns object: The new record + """ + return cls.create( + **{k: v for k, v in inst.__dict__.items() if k != "_sa_instance_state"} + ) + def update(self, **kwargs): """ Update an entry @@ -196,7 +210,6 @@ def update(self, **kwargs): self.save() return self - @classmethod def query(cls, *args): """ @@ -234,10 +247,12 @@ def delete(self, delete=True, hard_delete=False): self.db.rollback() raise + class Model(BaseModel): """ Model create """ + id = Column(Integer, primary_key=True) created_at = Column(sa_utils.ArrowType, default=utcnow) updated_at = Column(sa_utils.ArrowType, default=utcnow, onupdate=utcnow) @@ -270,9 +285,7 @@ def get(cls, id, include_deleted=False): :param id: The id of the entry :param include_deleted: It should not query deleted record. Set to True to get all """ - return cls.query(include_deleted=include_deleted)\ - .filter(cls.id == id)\ - .first() + return cls.query(include_deleted=include_deleted).filter(cls.id == id).first() def delete(self, delete=True, hard_delete=False): """ @@ -289,10 +302,7 @@ def delete(self, delete=True, hard_delete=False): self.db.rollback() raise else: - data = { - "is_deleted": delete, - "deleted_at": utcnow() if delete else None - } + data = {"is_deleted": delete, "deleted_at": utcnow() if delete else None} self.update(**data) return self @@ -334,14 +344,17 @@ class User(db.Model): """ - def __init__(self, uri='sqlite://', - app=None, - echo=False, - pool_size=None, - pool_timeout=None, - pool_recycle=None, - convert_unicode=True, - query_cls=BaseQuery): + def __init__( + self, + uri="sqlite://", + app=None, + echo=False, + pool_size=None, + pool_timeout=None, + pool_recycle=None, + query_cls=BaseQuery, + connect_args={}, + ): self.uri = uri self.info = make_url(uri) @@ -350,18 +363,21 @@ def __init__(self, uri='sqlite://', pool_size=pool_size, pool_timeout=pool_timeout, pool_recycle=pool_recycle, - convert_unicode=convert_unicode, + connect_args=connect_args, ) self.connector = None self._engine_lock = threading.Lock() self.session = _create_scoped_session(self, query_cls=query_cls) - self.Model = declarative_base(cls=Model, name='Model') - self.BaseModel = declarative_base(cls=BaseModel, name='BaseModel') + self.Model = declarative_base(cls=Model, name="Model") + self.BaseModel = declarative_base(cls=BaseModel, name="BaseModel") self.Model.db, self.BaseModel.db = self, self - self.Model._query, self.BaseModel._query = self.session.query, self.session.query + self.Model._query, self.BaseModel._query = ( + self.session.query, + self.session.query, + ) if app is not None: self.init_app(app) @@ -369,25 +385,21 @@ def __init__(self, uri='sqlite://', _include_sqlalchemy(self) def _cleanup_options(self, **kwargs): - options = dict([ - (key, val) - for key, val in kwargs.items() - if val is not None - ]) + options = dict([(key, val) for key, val in kwargs.items() if val is not None]) return self._apply_driver_hacks(options) def _apply_driver_hacks(self, options): if "mysql" in self.info.drivername: - self.info.query.setdefault('charset', 'utf8') - options.setdefault('pool_size', 10) - options.setdefault('pool_recycle', 7200) - elif self.info.drivername == 'sqlite': - no_pool = options.get('pool_size') == 0 - memory_based = self.info.database in (None, '', ':memory:') + self.info.query.setdefault("charset", "utf8") + options.setdefault("pool_size", 10) + options.setdefault("pool_recycle", 7200) + elif self.info.drivername == "sqlite": + no_pool = options.get("pool_size") == 0 + memory_based = self.info.database in (None, "", ":memory:") if memory_based and no_pool: raise ValueError( - 'SQLite in-memory database with an empty queue' - ' (pool_size = 0) is not possible due to data loss.' + "SQLite in-memory database with an empty queue" + " (pool_size = 0) is not possible due to data loss." ) return options @@ -397,7 +409,7 @@ def init_app(self, app): environment, never use a database without initialize it first, or connections will leak. """ - if not hasattr(app, 'databases'): + if not hasattr(app, "databases"): app.databases = [] if isinstance(app.databases, list): if self in app.databases: @@ -417,14 +429,14 @@ def rollback(error=None): self.set_flask_hooks(app, shutdown, rollback) def set_flask_hooks(self, app, shutdown, rollback): - if hasattr(app, 'after_request'): + if hasattr(app, "after_request"): app.after_request(shutdown) - if hasattr(app, 'on_exception'): + if hasattr(app, "on_exception"): app.on_exception(rollback) @property def engine(self): - """Gives access to the engine. """ + """Gives access to the engine.""" with self._engine_lock: connector = self.connector if connector is None: @@ -459,15 +471,15 @@ def rollback(self): return self.session.rollback() def create_all(self): - """Creates all tables. """ + """Creates all tables.""" self.Model.metadata.create_all(bind=self.engine) def drop_all(self): - """Drops all tables. """ + """Drops all tables.""" self.Model.metadata.drop_all(bind=self.engine) def reflect(self, meta=None): - """Reflects tables from the database. """ + """Reflects tables from the database.""" meta = meta or MetaData() meta.reflect(bind=self.engine) return meta