Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
114 changes: 63 additions & 51 deletions active_alchemy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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


Expand All @@ -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)
Expand All @@ -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):
"""
Expand Down Expand Up @@ -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
Expand All @@ -196,7 +210,6 @@ def update(self, **kwargs):
self.save()
return self


@classmethod
def query(cls, *args):
"""
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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):
"""
Expand All @@ -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

Expand Down Expand Up @@ -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)
Expand All @@ -350,44 +363,43 @@ 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)

_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

Expand All @@ -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:
Expand All @@ -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:
Expand Down Expand Up @@ -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
Expand Down