"""SQLAlchemy connector.""" from contextlib import contextmanager import functools from sqlalchemy import create_engine from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.orm import sessionmaker from availability import config BaseModel = declarative_base() _db_engine = create_engine( config.DB_CONNECTION_STRING, connect_args=config.DB_CONNECT_ARGS, pool_recycle=config.DB_CONNECTION_POOL_RECYCLE_TIMEOUT) _db_session_factory = sessionmaker(bind=_db_engine, expire_on_commit=False) @contextmanager def session_scope(): """Provide a transactional scope around a series of operations.""" session = _db_session_factory() try: yield session session.commit() except Exception: session.rollback() raise finally: session.close() def db_session_wrap(func): """Shared DB Session wrapper function. Creates a new session if one isn't passed in. This lets us share/pass-in a common session across multiple functions, making them all transactional. """ @functools.wraps(func) def wrapper(*args, **kwargs): if 'session' in kwargs: return func(*args, **kwargs) else: with session_scope() as session: return func(*args, session=session, **kwargs) return wrapper