"""Utility functions.""" import contextlib @contextlib.contextmanager def get_connection(engine, initial_statements=None): """Provide a transactional scope around a series of operations. Args: engine (sqlalchemy.engine.Engine): An engine to use for connection. initial_statements (iterable): An iterable of SQL statements to perform immediately after creating a connection. Yields: sqlalchemy.engine.Connection: An established connection. """ conn = engine.connect() trans = conn.begin() try: if initial_statements: for statement in initial_statements: conn.execute(statement) yield conn trans.commit() except: trans.rollback() raise finally: conn.close()