diff --git a/tests/test_database.py b/tests/test_database.py index c3aed6c..8becec3 100644 --- a/tests/test_database.py +++ b/tests/test_database.py @@ -16,24 +16,40 @@ def db_session() -> db.dal.Session: db.dal.engine.pool.recreate() +@fixture(scope='function') +def session(db_session) -> db.dal.Session: + """Provides a transactional scope around each test.""" + connection = db_session.connection() + transaction = connection.begin() + + # bind the session to this specific connection + db_session.bind = connection + + yield db_session + + # roll all changed that occurred back + transaction.rollback() + connection.close() + + def test_base() -> None: """Test to ensure db.Base() is a sqlalchemy DeclarativeBase object.""" assert isinstance(db.Base(), sqlalchemy.orm.DeclarativeBase) -def test_powerball_draw_by_date_blank(db_session) -> None: +def test_powerball_draw_by_date_blank(session) -> None: """ Test that a null record is returned when no date is provided and that no dateless records can exist in the table. """ - with db_session as session: + with session: blank_results = session.query(db.PowerballDraw).filter_by(draw_date='').all() assert blank_results == [] -def test_powerball_draw_by_date(db_session) -> None: +def test_powerball_draw_by_date(session) -> None: """Test to be sure we can find the first draw in the table by date.""" - with db_session as session: + with session: date = datetime(2010, 2, 3).date() draw = db.PowerballDraw(draw_date=date, main_ball1=17, main_ball2=22, main_ball3=36, main_ball4=37, main_ball5=52, powerball=24, power_play=2) @@ -42,9 +58,9 @@ def test_powerball_draw_by_date(db_session) -> None: assert record == draw -def test_powerball_draw_integrity(db_session) -> None: +def test_powerball_draw_integrity(session) -> None: """Test that a proper error is received when a duplicate value insertion is attempted.""" - with db_session as session: + with session: date = datetime.strptime('02/03/2010','%m/%d/%Y').date() draw = db.PowerballDraw(draw_date=date, main_ball1=17, main_ball2=22, main_ball3=36, main_ball4=37, main_ball5=52, powerball=24, power_play=2) @@ -59,9 +75,9 @@ def test_powerball_draw_integrity(db_session) -> None: session.close() -def test_powerball_repr_and_str(db_session) -> None: +def test_powerball_repr_and_str(session) -> None: """Test the __repr__ and __str__ methods of PowerballDraw objects.""" - with db_session as session: + with session: date = datetime(2010, 2, 3).date() record = session.query(db.PowerballDraw).filter_by(draw_date=date).first() assert str(record) == '2010-02-03, [17, 22, 36, 37, 52], 24, 2x'