diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..25ebbc1 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,57 @@ +# tests/conftest.py +import datetime +from typing import Any, Generator + +import pytest +from sqlalchemy import create_engine, Date, Engine, Integer +from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, sessionmaker, Session + +class Base(DeclarativeBase): + pass + + +class PowerballDraw(Base): + __tablename__ = 'powerball_draws' + + draw_date: Mapped[datetime.date] = mapped_column(Date, primary_key=True) + main_ball1: Mapped[int] = mapped_column(Integer) + main_ball2: Mapped[int] = mapped_column(Integer) + main_ball3: Mapped[int] = mapped_column(Integer) + main_ball4: Mapped[int] = mapped_column(Integer) + main_ball5: Mapped[int] = mapped_column(Integer) + powerball: Mapped[int] = mapped_column(Integer) + power_play: Mapped[int] = mapped_column(Integer) + + +@pytest.fixture(scope='session') +def db_engine() -> Generator[Engine, Any, None]: + # create a database connection engine + engine = create_engine('sqlite:///:memory:') + + # create all the objects defined in the Base + Base.metadata.create_all(engine) + # yield the engine + yield engine + # drop all objects to clean up + Base.metadata.drop_all(engine) + + +@pytest.fixture(scope='function') +def db_session(db_engine) -> Generator[Session, Any, None]: + # create a database connection + connection = db_engine.connect() + + # start a transaction + transaction = connection.begin() + + # create a session object + db_session = sessionmaker(bind=connection) + session = db_session() + + # yield the session object + yield session + + # tear the session down and close the connection + session.close() + transaction.rollback() + connection.close()