Skip to content

Instantly share code, notes, and snippets.

@portothree
Last active September 6, 2021 18:15
Show Gist options
  • Select an option

  • Save portothree/fb092482eebd75f0731c3877d682617c to your computer and use it in GitHub Desktop.

Select an option

Save portothree/fb092482eebd75f0731c3877d682617c to your computer and use it in GitHub Desktop.
Sqlalchemy's session injection with dependency-injector lib
import contextlib
import logging
import dependency_injector.containers as containers
import dependency_injector.providers as providers
import sqlalchemy
from sqlalchemy import orm, MetaData, Column, Integer, String, Table
logger = logging.getLogger(__name__)
class User:
"""Example user entity."""
def __init__(self, name):
self.name = name
def __repr__(self):
return f'User(name="{self.name}")'
class Database:
"""Acts as a connection pool and session factory."""
def __init__(self, db_url):
self.db_url = db_url
self.engine = sqlalchemy.create_engine(db_url)
self.session_factory = orm.scoped_session(orm.sessionmaker(bind=self.engine))
def prepare_db(self):
# Prepare mappings
metadata = MetaData()
user_table = Table(
'user',
metadata,
Column('id', Integer, primary_key=True),
Column('name', String(50)),
)
orm.mapper(User, user_table)
# Create database / tables
metadata.create_all(bind=self.engine)
# Setup test data
user = User(name='asyncee')
with self.session() as s:
s.add(user)
s.commit()
@contextlib.contextmanager
def session(self):
session = self.session_factory()
try:
yield session
except Exception as e:
logger.error('Session rollback because of exception: %s', e, exc_info=True)
session.rollback()
finally:
session.close()
@staticmethod
def session_factory(db):
return db.session
class UserRepository:
def __init__(self, session):
self.session = session
def get_users(self):
return self.session.query(User).all()
class UsersService:
def __init__(self, repository):
self.repository = repository
def execute(self):
return self.get_users()
def get_users(self):
return self.repository.get_users()
class PrintUsersUseCase:
"""Print all users in the system to stdout."""
def __init__(self, session_factory, repository_class):
self.session_factory = session_factory
self.repository_class = repository_class
def run(self):
with self.session_factory() as session:
# Repository or service must know only about session, because
# many components may run in one transaction scope.
repository = self.repository_class(session)
users_service = UsersService(repository)
users = users_service.execute()
# This is business logic of this use-case.
print(users)
# Persist all changes into database if needed.
session.commit()
class Core(containers.DeclarativeContainer):
config = providers.Configuration('app_config')
class Sqla(containers.DeclarativeContainer):
db = providers.Singleton(Database, db_url=Core.config.db_url)
session_factory = providers.Factory(Database.session_factory, db=db)
class UseCases(containers.DeclarativeContainer):
get_user = providers.Factory(
GetUsersUseCase, repository_class=UserRepository,
session_factory=Sqla.session_factory)
if __name__ == "__main__":
Core.config.override({'db_url': 'sqlite:///:memory:'})
Sqla.db().prepare_db()
use_case = UseCases.get_user()
use_case.run()
dependency-injector==4.36.0
SQLAlchemy==1.4.23
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment